From 3d16e216fe023bd93d33ffc78654b638bcf04859 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 23 Jul 2026 15:12:57 +0000 Subject: [PATCH 01/14] feat(code): integrate Hooks v2 server lifecycle events Wire PreToolUse, PostToolUse, Stop, SubagentStart, and SubagentStop through LangGraph interrupts so the client HooksRuntime can execute handlers and return typed decisions. Co-authored-by: Johannes du Plessis --- libs/code/deepagents_code/_cli_context.py | 23 + libs/code/deepagents_code/agent.py | 8 + libs/code/deepagents_code/app.py | 14 +- .../deepagents_code/client/non_interactive.py | 58 +- libs/code/deepagents_code/hooks/client.py | 67 ++ libs/code/deepagents_code/hooks/context.py | 38 ++ libs/code/deepagents_code/hooks/interrupt.py | 130 ++++ libs/code/deepagents_code/hooks/runtime.py | 10 + .../hooks/server_middleware.py | 593 ++++++++++++++++++ libs/code/deepagents_code/hooks/snapshot.py | 16 + .../deepagents_code/tui/textual_adapter.py | 40 +- .../unit_tests/hooks/test_server_lifecycle.py | 191 ++++++ 12 files changed, 1185 insertions(+), 3 deletions(-) create mode 100644 libs/code/deepagents_code/hooks/client.py create mode 100644 libs/code/deepagents_code/hooks/context.py create mode 100644 libs/code/deepagents_code/hooks/interrupt.py create mode 100644 libs/code/deepagents_code/hooks/server_middleware.py create mode 100644 libs/code/tests/unit_tests/hooks/test_server_lifecycle.py diff --git a/libs/code/deepagents_code/_cli_context.py b/libs/code/deepagents_code/_cli_context.py index b00f6d6896..0738c6e8e3 100644 --- a/libs/code/deepagents_code/_cli_context.py +++ b/libs/code/deepagents_code/_cli_context.py @@ -51,6 +51,12 @@ class CLIContextSchema: offload_tool_call_id: str | None = None + hooks_snapshot_id: str | None = None + + hooks_server_events: list[str] = field(default_factory=list) + + prompt_id: str | None = None + class CLIContext(TypedDict, total=False): """Client-facing builder for the per-run graph context payload. @@ -107,3 +113,20 @@ class CLIContext(TypedDict, total=False): This is set by the client, not graph state, so model-generated calls cannot grant themselves permission to execute during the hidden compaction turn. """ + + hooks_snapshot_id: str | None + """Canonical Hooks v2 configuration hash for this session. + + Server-owned lifecycle middleware includes this id on interrupt requests so + the client can reject mismatched resumes. + """ + + hooks_server_events: list[str] + """Server-owned HookEvent names that have configured handlers. + + Middleware only interrupts for events listed here, avoiding a round-trip + when the session snapshot has no matching handlers. + """ + + prompt_id: str | None + """Optional per-turn prompt id projected into hook context.""" diff --git a/libs/code/deepagents_code/agent.py b/libs/code/deepagents_code/agent.py index f1a92eb2eb..6c20c1c992 100644 --- a/libs/code/deepagents_code/agent.py +++ b/libs/code/deepagents_code/agent.py @@ -2713,6 +2713,14 @@ def _subagent_cli_middleware( if restrictive_shell_allow_list is not None: agent_middleware.append(ShellAllowListMiddleware(restrictive_shell_allow_list)) + # Server-owned Hooks v2 lifecycle events (Pre/Post tool, Stop, subagent). + # Gated at runtime by `hooks_server_events` on the per-run context so idle + # sessions without configured handlers pay no interrupt round-trip. + from deepagents_code.hooks.server_middleware import ServerHooksMiddleware + + hooks_cwd = Path(effective_cwd) if effective_cwd is not None else Path.cwd() + agent_middleware.append(ServerHooksMiddleware(cwd=hooks_cwd)) + # Get or use custom system prompt if system_prompt is None: system_prompt = get_system_prompt( diff --git a/libs/code/deepagents_code/app.py b/libs/code/deepagents_code/app.py index b89b4d19c0..6b90a2fce4 100644 --- a/libs/code/deepagents_code/app.py +++ b/libs/code/deepagents_code/app.py @@ -2261,6 +2261,8 @@ def __init__( # Assign the backing field directly: the setter reads `self._thread_id` # to detect a thread change, and it isn't set yet. self._thread_id = thread_id or _new_thread_id() + self.hooks_runtime = None + """Optional session-scoped Hooks v2 client runtime.""" @property def auto_approve(self) -> bool: @@ -4279,10 +4281,20 @@ async def _init_session_state(self) -> None: """Create session state in a thread (imports deepagents_code.sessions).""" def _create() -> TextualSessionState: - return TextualSessionState( + from pathlib import Path + + from deepagents_code.hooks.runtime import HooksRuntime + + state = TextualSessionState( approval_mode=self._approval_mode, thread_id=self._lc_thread_id, ) + try: + state.hooks_runtime = HooksRuntime.create(cwd=Path(self._cwd)) + except Exception: + logger.exception("Failed to create HooksRuntime; server hooks disabled") + state.hooks_runtime = None + return state try: session_state = await asyncio.to_thread(_create) diff --git a/libs/code/deepagents_code/client/non_interactive.py b/libs/code/deepagents_code/client/non_interactive.py index 4f11d56747..9f05e77ab5 100644 --- a/libs/code/deepagents_code/client/non_interactive.py +++ b/libs/code/deepagents_code/client/non_interactive.py @@ -354,6 +354,15 @@ class StreamState: Used to resume the agent after HITL processing. """ + pending_hook_interrupts: dict[str, object] = field(default_factory=dict) + """Raw Hooks v2 invocation interrupt payloads awaiting client fulfillment.""" + + hook_response: dict[str, Any] = field(default_factory=dict) + """Resume values for fulfilled Hooks v2 interrupts, keyed by interrupt id.""" + + hooks_runtime: Any | None = None + """Optional session-scoped HooksRuntime used to fulfill server hook interrupts.""" + interrupt_occurred: bool = False """Flag indicating whether any HITL interrupt was received during the current stream pass.""" @@ -419,9 +428,15 @@ def _process_interrupts( state: Stream state to update with new pending interrupts. console: Rich console for user-visible warnings. """ + from deepagents_code.hooks.interrupt import is_hook_interrupt_payload + interrupts = data["__interrupt__"] if interrupts: for interrupt_obj in interrupts: + if is_hook_interrupt_payload(interrupt_obj.value): + state.pending_hook_interrupts[interrupt_obj.id] = interrupt_obj.value + state.interrupt_occurred = True + continue try: validated_request = _HITL_REQUEST_ADAPTER.validate_python( interrupt_obj.value @@ -933,6 +948,30 @@ def _collect_action_request_warnings(action_request: ActionRequest) -> list[str] return warnings +async def _fulfill_pending_hook_interrupts(state: StreamState) -> None: + """Execute pending server-owned hook interrupts on the client runtime. + + Raises: + RuntimeError: If a hook interrupt arrives without a session runtime, or + if a payload cannot be parsed. + """ + if not state.pending_hook_interrupts: + return + from deepagents_code.hooks.client import fulfill_hook_interrupt + + if state.hooks_runtime is None: + msg = "Received hook invocation interrupt without a HooksRuntime" + raise RuntimeError(msg) + pending = dict(state.pending_hook_interrupts) + state.pending_hook_interrupts.clear() + for interrupt_id, payload in pending.items(): + resume_value = await fulfill_hook_interrupt(state.hooks_runtime, payload) + if resume_value is None: + msg = f"Failed to parse hook interrupt {interrupt_id}" + raise RuntimeError(msg) + state.hook_response[interrupt_id] = resume_value + + def _process_hitl_interrupts(state: StreamState, console: Console) -> None: """Iterate over pending HITL interrupts and build approval/rejection responses. @@ -1103,6 +1142,20 @@ async def _run_agent_loop( # unset in context rather than passing a blank string to model middleware. context_thread_id = thread_id if isinstance(thread_id, str) and thread_id else None context = CLIContext(thread_id=context_thread_id) + + from pathlib import Path + + from deepagents_code.hooks.context import apply_hooks_context + from deepagents_code.hooks.runtime import HooksRuntime + + try: + hooks_runtime = HooksRuntime.create(cwd=Path.cwd()) + except Exception: + logger.exception("Failed to create HooksRuntime; server hooks disabled") + hooks_runtime = None + apply_hooks_context(context, hooks_runtime) + state.hooks_runtime = hooks_runtime + await dispatch_hook("session.start", {"thread_id": thread_id}) start_time = time.monotonic() @@ -1137,8 +1190,11 @@ async def _run_agent_loop( turns += 1 state.interrupt_occurred = False state.hitl_response.clear() + state.hook_response.clear() + await _fulfill_pending_hook_interrupts(state) _process_hitl_interrupts(state, console) - stream_input = Command(resume=state.hitl_response) + resume_payload = {**state.hook_response, **state.hitl_response} + stream_input = Command(resume=resume_payload) await _stream_agent( agent, stream_input, config, state, console, file_op_tracker, context ) diff --git a/libs/code/deepagents_code/hooks/client.py b/libs/code/deepagents_code/hooks/client.py new file mode 100644 index 0000000000..b98880cb70 --- /dev/null +++ b/libs/code/deepagents_code/hooks/client.py @@ -0,0 +1,67 @@ +"""Client-side fulfillment for server-owned Hooks v2 interrupts.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +from deepagents_code.hooks.interrupt import ( + build_hook_resume_value, + parse_hook_interrupt_payload, +) +from deepagents_code.hooks.models.transport import HookInvocationResponse + +if TYPE_CHECKING: + from deepagents_code.hooks.models.transport import HookInvocationRequest + from deepagents_code.hooks.runtime import HooksRuntime + + +async def fulfill_hook_invocation( + runtime: HooksRuntime, + request: HookInvocationRequest, +) -> dict[str, object]: + """Execute a server-owned hook request and return a resume payload. + + Args: + runtime: Session-scoped client Hooks runtime. + request: Validated invocation request from the server. + + Returns: + JSON-compatible resume value for `Command(resume=...)`. + + Raises: + ValueError: If the request snapshot does not match this session. + """ + if request.snapshot_id != runtime.snapshot_id: + msg = ( + f"Hook snapshot mismatch: request {request.snapshot_id} != " + f"runtime {runtime.snapshot_id}" + ) + raise ValueError(msg) + + decision = await runtime.invoke(request.invocation) + response = HookInvocationResponse( + protocol_version=1, + invocation_id=request.invocation_id, + snapshot_id=request.snapshot_id, + decision=decision, + ) + return build_hook_resume_value(response) + + +async def fulfill_hook_interrupt( + runtime: HooksRuntime, + interrupt_value: object, +) -> dict[str, object] | None: + """Fulfill a raw interrupt value when it is a hook invocation. + + Args: + runtime: Session-scoped client Hooks runtime. + interrupt_value: Raw LangGraph interrupt payload. + + Returns: + Resume value for hook interrupts, otherwise `None`. + """ + request = parse_hook_interrupt_payload(interrupt_value) + if request is None: + return None + return await fulfill_hook_invocation(runtime, request) diff --git a/libs/code/deepagents_code/hooks/context.py b/libs/code/deepagents_code/hooks/context.py new file mode 100644 index 0000000000..5b353c0143 --- /dev/null +++ b/libs/code/deepagents_code/hooks/context.py @@ -0,0 +1,38 @@ +"""Helpers for attaching Hooks v2 session identity to graph context.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from deepagents_code._cli_context import CLIContext + from deepagents_code.hooks.runtime import HooksRuntime + + +def apply_hooks_context( + context: CLIContext, + runtime: HooksRuntime | None, + *, + prompt_id: str | None = None, +) -> CLIContext: + """Attach Hooks v2 snapshot identity and server event gates to `context`. + + Args: + context: Mutable per-run graph context. + runtime: Session Hooks runtime, or `None` when hooks are unavailable. + prompt_id: Optional per-turn prompt id. + + Returns: + The same context mapping, updated in place. + """ + if runtime is None: + context.pop("hooks_snapshot_id", None) + context.pop("hooks_server_events", None) + else: + context["hooks_snapshot_id"] = runtime.snapshot_id + context["hooks_server_events"] = list(runtime.configured_server_events()) + if prompt_id is not None: + context["prompt_id"] = prompt_id + else: + context.pop("prompt_id", None) + return context diff --git a/libs/code/deepagents_code/hooks/interrupt.py b/libs/code/deepagents_code/hooks/interrupt.py new file mode 100644 index 0000000000..aaba8e94ac --- /dev/null +++ b/libs/code/deepagents_code/hooks/interrupt.py @@ -0,0 +1,130 @@ +"""Client↔server interrupt transport for Hooks v2 server-owned events.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any, Literal, TypeAlias + +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter + +from deepagents_code.hooks.models.adapters import ( + HOOK_INVOCATION_RESPONSE_ADAPTER, +) +from deepagents_code.hooks.models.transport import ( # noqa: TC001 - Pydantic runtime + HookInvocationRequest, + HookInvocationResponse, +) + +if TYPE_CHECKING: + from uuid import UUID + +HOOK_INVOCATION_INTERRUPT_TYPE: Literal["hook_invocation"] = "hook_invocation" + + +class HookInvocationInterrupt(BaseModel): + """LangGraph interrupt envelope for a server-owned hook invocation.""" + + model_config = ConfigDict(extra="forbid") + + type: Literal["hook_invocation"] = HOOK_INVOCATION_INTERRUPT_TYPE + request: HookInvocationRequest + + +HOOK_INVOCATION_INTERRUPT_ADAPTER = TypeAdapter(HookInvocationInterrupt) + +HookResumeValue: TypeAlias = dict[str, Any] + + +def build_hook_interrupt_payload(request: HookInvocationRequest) -> dict[str, Any]: + """Serialize a hook invocation request for `interrupt()`. + + Args: + request: Versioned server-owned invocation request. + + Returns: + JSON-compatible interrupt payload with a stable `type` discriminator. + """ + return HOOK_INVOCATION_INTERRUPT_ADAPTER.dump_python( + HookInvocationInterrupt(request=request), + mode="json", + ) + + +def parse_hook_interrupt_payload(value: object) -> HookInvocationRequest | None: + """Parse a hook invocation interrupt payload when present. + + Args: + value: Raw interrupt value from LangGraph. + + Returns: + The embedded request, or `None` when `value` is not a hook interrupt. + """ + if ( + not isinstance(value, dict) + or value.get("type") != HOOK_INVOCATION_INTERRUPT_TYPE + ): + return None + interrupt = HOOK_INVOCATION_INTERRUPT_ADAPTER.validate_python(value) + return interrupt.request + + +def build_hook_resume_value(response: HookInvocationResponse) -> HookResumeValue: + """Serialize a hook invocation response for `Command(resume=...)`. + + Args: + response: Decision returned by the client runtime. + + Returns: + JSON-compatible resume value for the matching interrupt id. + """ + return HOOK_INVOCATION_RESPONSE_ADAPTER.dump_python(response, mode="json") + + +def parse_hook_resume_value( + value: object, + *, + invocation_id: UUID, + snapshot_id: str, +) -> HookInvocationResponse: + """Validate a resumed hook response against the outstanding request. + + Args: + value: Resume payload returned by the client. + invocation_id: Expected invocation id from the request. + snapshot_id: Expected configuration snapshot id. + + Returns: + Validated response. + + Raises: + ValueError: If the resume payload is missing, mistyped, or mismatched. + """ + response = HOOK_INVOCATION_RESPONSE_ADAPTER.validate_python(value) + if response.invocation_id != invocation_id: + msg = ( + f"Hook resume invocation_id mismatch: expected {invocation_id}, " + f"got {response.invocation_id}" + ) + raise ValueError(msg) + if response.snapshot_id != snapshot_id: + msg = ( + f"Hook resume snapshot_id mismatch: expected {snapshot_id}, " + f"got {response.snapshot_id}" + ) + raise ValueError(msg) + return response + + +def is_hook_interrupt_payload(value: object) -> bool: + """Return whether `value` looks like a Hooks v2 invocation interrupt.""" + return ( + isinstance(value, dict) and value.get("type") == HOOK_INVOCATION_INTERRUPT_TYPE + ) + + +class HookInterruptMismatch(BaseModel): + """Diagnostic retained when a hook resume cannot be applied.""" + + model_config = ConfigDict(extra="forbid") + + code: Literal["hook_resume_mismatch"] = "hook_resume_mismatch" + message: str = Field(min_length=1) diff --git a/libs/code/deepagents_code/hooks/runtime.py b/libs/code/deepagents_code/hooks/runtime.py index 84e42d07c0..c383866191 100644 --- a/libs/code/deepagents_code/hooks/runtime.py +++ b/libs/code/deepagents_code/hooks/runtime.py @@ -92,6 +92,16 @@ def snapshot_id(self) -> str: """Canonical configuration hash for this session.""" return self.snapshot.snapshot_id + def configured_server_events(self) -> tuple[str, ...]: + """Stable event names the server should emit for this session. + + Returns: + Sorted HookEvent values that have configured server-owned handlers. + """ + return tuple( + sorted(event.value for event in self.snapshot.configured_server_events()) + ) + def append_messages( self, thread_id: str, diff --git a/libs/code/deepagents_code/hooks/server_middleware.py b/libs/code/deepagents_code/hooks/server_middleware.py new file mode 100644 index 0000000000..e56c7423ee --- /dev/null +++ b/libs/code/deepagents_code/hooks/server_middleware.py @@ -0,0 +1,593 @@ +"""Server-owned Hooks v2 lifecycle middleware. + +Emits `PreToolUse`, `PostToolUse`, `Stop`, `SubagentStart`, and `SubagentStop` +through the LangGraph interrupt channel so the client runtime can execute +matching handlers and return typed decisions. +""" + +from __future__ import annotations + +import time +from collections.abc import Mapping, Sequence +from datetime import UTC, datetime, timedelta +from typing import TYPE_CHECKING, Any, NotRequired, cast +from uuid import UUID, uuid4 + +from langchain.agents.middleware.types import ( + AgentMiddleware, + AgentState, + ContextT, + ResponseT, + hook_config, +) +from langchain_core.messages import AIMessage, HumanMessage, ToolMessage +from langgraph.types import Command, interrupt +from typing_extensions import TypedDict + +from deepagents_code.approval_mode import ApprovalMode, coerce_approval_mode +from deepagents_code.hooks.interrupt import ( + build_hook_interrupt_payload, + parse_hook_resume_value, +) +from deepagents_code.hooks.models.domain import ( + AgentIdentity, + HookContext, + HookDecision, + HookEvent, + HookInvocation, + PermissionEffect, + PostToolUseDecision, + PostToolUseEvent, + PreToolUseDecision, + PreToolUseEvent, + StopDecision, + StopEvent, + SubagentStartDecision, + SubagentStartEvent, + SubagentStopDecision, + SubagentStopEvent, + ToolCallData, +) +from deepagents_code.hooks.models.transport import HookInvocationRequest +from deepagents_code.hooks.tools import to_wire_tool_name + +if TYPE_CHECKING: + from collections.abc import Awaitable, Callable + from pathlib import Path + + from langchain.tools.tool_node import ToolCallRequest + from langchain_core.messages.tool import ToolCall + from langgraph.runtime import Runtime + + from deepagents_code.json_types import JsonObject + +_DEFAULT_DEADLINE = timedelta(seconds=600) +_STOP_STATE_KEY = "_hooks_stop_continuation_count" +_TASK_TOOL_NAME = "task" + + +class ServerHooksState(AgentState[Any]): + """Agent state extensions for server-owned hook middleware.""" + + _hooks_stop_continuation_count: NotRequired[int] + + +class _SessionHookGate(TypedDict): + snapshot_id: str + events: frozenset[str] + + +class ServerHooksMiddleware(AgentMiddleware[ServerHooksState, ContextT, ResponseT]): + """Emit server-owned lifecycle events over the hook interrupt transport.""" + + state_schema = ServerHooksState + + def __init__( + self, + *, + cwd: Path, + default_deadline: timedelta = _DEFAULT_DEADLINE, + ) -> None: + """Initialize middleware. + + Args: + cwd: Session working directory projected into hook context. + default_deadline: Client execution deadline attached to requests. + """ + super().__init__() + self._cwd = cwd + self._default_deadline = default_deadline + + def wrap_tool_call( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], ToolMessage | Command[Any]], + ) -> ToolMessage | Command[Any]: + """Run Pre/Post tool hooks around a synchronous tool call. + + Returns: + Tool result, possibly rewritten by hook decisions. + """ + return self._run_tool_call(request, handler, async_handler=False) + + async def awrap_tool_call( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command[Any]]], + ) -> ToolMessage | Command[Any]: + """Run Pre/Post tool hooks around an asynchronous tool call. + + Returns: + Tool result, possibly rewritten by hook decisions. + """ + return await self._run_tool_call_async(request, handler) + + @hook_config(can_jump_to=["model"]) + def after_agent( + self, + state: ServerHooksState, + runtime: Runtime[ContextT], + ) -> dict[str, Any] | None: + """Emit `Stop` when the agent reaches a natural end. + + Returns: + Optional state update that may jump back to the model. + """ + return self._after_agent(state, runtime) + + @hook_config(can_jump_to=["model"]) + async def aafter_agent( + self, + state: ServerHooksState, + runtime: Runtime[ContextT], + ) -> dict[str, Any] | None: + """Async `Stop` emission; mirrors `after_agent`. + + Returns: + Optional state update that may jump back to the model. + """ + return self._after_agent(state, runtime) + + def _run_tool_call( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], ToolMessage | Command[Any]], + *, + async_handler: bool, + ) -> ToolMessage | Command[Any]: + del async_handler + gate = _session_gate(request.runtime.context) + call = _tool_call_data(request) + context = _hook_context( + request.runtime.context, request.runtime.config, self._cwd + ) + request = self._maybe_subagent_start(request, call, context, gate) + denied = self._maybe_pre_tool_use(call, context, gate, request.runtime.config) + if denied is not None: + return denied + started = time.perf_counter() + result = handler(request) + duration_ms = int((time.perf_counter() - started) * 1000) + result = self._maybe_post_tool_use( + call, context, gate, request.runtime.config, result, duration_ms + ) + return self._maybe_subagent_stop( + call, context, gate, request.runtime.config, result + ) + + async def _run_tool_call_async( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command[Any]]], + ) -> ToolMessage | Command[Any]: + gate = _session_gate(request.runtime.context) + call = _tool_call_data(request) + context = _hook_context( + request.runtime.context, request.runtime.config, self._cwd + ) + request = self._maybe_subagent_start(request, call, context, gate) + denied = self._maybe_pre_tool_use(call, context, gate, request.runtime.config) + if denied is not None: + return denied + started = time.perf_counter() + result = await handler(request) + duration_ms = int((time.perf_counter() - started) * 1000) + result = self._maybe_post_tool_use( + call, context, gate, request.runtime.config, result, duration_ms + ) + return self._maybe_subagent_stop( + call, context, gate, request.runtime.config, result + ) + + def _maybe_subagent_start( + self, + request: ToolCallRequest, + call: ToolCallData, + context: HookContext, + gate: _SessionHookGate | None, + ) -> ToolCallRequest: + if call.name != _TASK_TOOL_NAME or not _event_enabled( + gate, HookEvent.SUBAGENT_START + ): + return request + agent = _task_agent_identity(call) + decision = _invoke_hook( + context, + SubagentStartEvent(event=HookEvent.SUBAGENT_START, agent=agent), + gate=gate, + config=request.runtime.config, + deadline=self._default_deadline, + ) + if not isinstance(decision, SubagentStartDecision): + msg = "Expected SubagentStartDecision, got {type(decision).__name__}" + raise TypeError(msg) + return _inject_subagent_start_context(request, decision) + + def _maybe_pre_tool_use( + self, + call: ToolCallData, + context: HookContext, + gate: _SessionHookGate | None, + config: Mapping[str, Any] | None, + ) -> ToolMessage | None: + if not _event_enabled(gate, HookEvent.PRE_TOOL_USE): + return None + decision = _invoke_hook( + context, + PreToolUseEvent(event=HookEvent.PRE_TOOL_USE, call=call), + gate=gate, + config=config, + deadline=self._default_deadline, + ) + if not isinstance(decision, PreToolUseDecision): + msg = "Expected PreToolUseDecision, got {type(decision).__name__}" + raise TypeError(msg) + return _denied_tool_message(call, decision.permission) + + def _maybe_post_tool_use( + self, + call: ToolCallData, + context: HookContext, + gate: _SessionHookGate | None, + config: Mapping[str, Any] | None, + result: ToolMessage | Command[Any], + duration_ms: int, + ) -> ToolMessage | Command[Any]: + if not _event_enabled(gate, HookEvent.POST_TOOL_USE): + return result + if not isinstance(result, ToolMessage): + return result + decision = _invoke_hook( + context, + PostToolUseEvent( + event=HookEvent.POST_TOOL_USE, + call=call, + result=result, + duration_ms=duration_ms, + ), + gate=gate, + config=config, + deadline=self._default_deadline, + ) + if not isinstance(decision, PostToolUseDecision): + msg = "Expected PostToolUseDecision, got {type(decision).__name__}" + raise TypeError(msg) + return _apply_post_tool_use(result, decision) + + def _maybe_subagent_stop( + self, + call: ToolCallData, + context: HookContext, + gate: _SessionHookGate | None, + config: Mapping[str, Any] | None, + result: ToolMessage | Command[Any], + ) -> ToolMessage | Command[Any]: + if call.name != _TASK_TOOL_NAME or not _event_enabled( + gate, HookEvent.SUBAGENT_STOP + ): + return result + agent = _task_agent_identity(call) + decision = _invoke_hook( + context, + SubagentStopEvent( + event=HookEvent.SUBAGENT_STOP, + agent=agent, + continuation_count=0, + last_assistant_message=_tool_result_text(result), + ), + gate=gate, + config=config, + deadline=self._default_deadline, + ) + if not isinstance(decision, SubagentStopDecision): + msg = "Expected SubagentStopDecision, got {type(decision).__name__}" + raise TypeError(msg) + return _apply_subagent_stop(result, decision) + + def _after_agent( + self, + state: ServerHooksState, + runtime: Runtime[ContextT], + ) -> dict[str, Any] | None: + gate = _session_gate(runtime.context) + if not _event_enabled(gate, HookEvent.STOP): + return None + continuation = int(state.get(_STOP_STATE_KEY, 0) or 0) + context = _hook_context(runtime.context, None, self._cwd) + decision = _invoke_hook( + context, + StopEvent( + event=HookEvent.STOP, + continuation_count=continuation, + last_assistant_message=_last_assistant_text(state.get("messages", ())), + ), + gate=gate, + config=None, + deadline=self._default_deadline, + ) + if not isinstance(decision, StopDecision): + msg = "Expected StopDecision, got {type(decision).__name__}" + raise TypeError(msg) + if not decision.continue_loop: + return None + feedback = "\n".join(decision.feedback).strip() or ( + decision.stop_reason or "Continue working." + ) + return { + "messages": [HumanMessage(content=feedback)], + "jump_to": "model", + _STOP_STATE_KEY: continuation + 1, + } + + +def _session_gate(runtime_context: object) -> _SessionHookGate | None: + fields = _context_mapping(runtime_context) + snapshot_id = fields.get("hooks_snapshot_id") + events = fields.get("hooks_server_events") + if not isinstance(snapshot_id, str) or not snapshot_id: + return None + if not isinstance(events, list) or not events: + return None + return { + "snapshot_id": snapshot_id, + "events": frozenset(str(item) for item in events), + } + + +def _event_enabled(gate: _SessionHookGate | None, event: HookEvent) -> bool: + return gate is not None and event.value in gate["events"] + + +def _invoke_hook( + context: HookContext, + event: ( + PreToolUseEvent + | PostToolUseEvent + | StopEvent + | SubagentStartEvent + | SubagentStopEvent + ), + *, + gate: _SessionHookGate | None, + config: Mapping[str, Any] | None, + deadline: timedelta, +) -> HookDecision: + if gate is None: + msg = "hooks_snapshot_id is required to emit server-owned hook events" + raise RuntimeError(msg) + request = HookInvocationRequest( + protocol_version=1, + invocation_id=uuid4(), + snapshot_id=gate["snapshot_id"], + run_id=_run_id(config), + invocation=HookInvocation(context=context, event=event), + deadline=datetime.now(UTC) + deadline, + ) + raw = interrupt(build_hook_interrupt_payload(request)) + response = parse_hook_resume_value( + raw, + invocation_id=request.invocation_id, + snapshot_id=request.snapshot_id, + ) + return response.decision + + +def _hook_context( + runtime_context: object, + config: Mapping[str, Any] | None, + cwd: Path, +) -> HookContext: + fields = _context_mapping(runtime_context) + thread_id = fields.get("thread_id") or _config_thread_id(config) or "unknown" + if not isinstance(thread_id, str): + thread_id = "unknown" + approval = coerce_approval_mode(fields.get("approval_mode", "manual")) + prompt_raw = fields.get("prompt_id") + prompt_id = UUID(prompt_raw) if isinstance(prompt_raw, str) and prompt_raw else None + return HookContext( + thread_id=thread_id, + cwd=cwd, + prompt_id=prompt_id, + approval_mode=( + approval if isinstance(approval, ApprovalMode) else ApprovalMode.MANUAL + ), + ) + + +def _context_mapping(runtime_context: object) -> dict[str, Any]: + if runtime_context is None: + return {} + if isinstance(runtime_context, Mapping): + return {str(key): value for key, value in runtime_context.items()} + result: dict[str, Any] = {} + for key in ( + "hooks_snapshot_id", + "hooks_server_events", + "thread_id", + "approval_mode", + "prompt_id", + ): + value = getattr(runtime_context, key, None) + if value is not None: + result[key] = value + return result + + +def _run_id(config: Mapping[str, Any] | None) -> str: + if isinstance(config, Mapping): + configurable = config.get("configurable") + if isinstance(configurable, Mapping): + for key in ("run_id", "thread_id"): + value = configurable.get(key) + if isinstance(value, str) and value: + return value + return str(uuid4()) + + +def _config_thread_id(config: Mapping[str, Any] | None) -> str | None: + if not isinstance(config, Mapping): + return None + configurable = config.get("configurable") + if not isinstance(configurable, Mapping): + return None + value = configurable.get("thread_id") + return value if isinstance(value, str) and value else None + + +def _tool_call_data(request: ToolCallRequest) -> ToolCallData: + tool_call = request.tool_call + raw_args = tool_call.get("args") + args: dict[str, Any] + if isinstance(raw_args, dict): + args = {str(key): value for key, value in raw_args.items()} + else: + args = {} + return ToolCallData( + id=str(tool_call.get("id") or ""), + name=str(tool_call.get("name") or ""), + args=cast("JsonObject", args), + mcp_server=_mcp_server_from_tool(request.tool), + ) + + +def _mcp_server_from_tool(tool: object | None) -> str | None: + if tool is None: + return None + metadata = getattr(tool, "metadata", None) + if not isinstance(metadata, Mapping): + return None + for key in ("mcp_server", "mcp_server_name", "server_name"): + value = metadata.get(key) + if isinstance(value, str) and value: + return value + return None + + +def _denied_tool_message( + call: ToolCallData, + permission: PermissionEffect, +) -> ToolMessage | None: + if permission.behavior not in {"deny", "ask"}: + return None + reason = permission.reason or ( + "Blocked by PreToolUse hook" + if permission.behavior == "deny" + else "PreToolUse hook requested approval (ask is not applied yet)" + ) + wire_name = to_wire_tool_name(call.name, mcp_server=call.mcp_server) + return ToolMessage( + content=f"{wire_name} blocked by hook: {reason}", + name=call.name, + tool_call_id=call.id, + status="error", + ) + + +def _apply_post_tool_use( + result: ToolMessage, + decision: PostToolUseDecision, +) -> ToolMessage: + extras: list[str] = [] + if decision.feedback: + extras.append("\n".join(decision.feedback)) + if decision.context: + extras.append("\n".join(decision.context)) + if not extras: + return result + suffix = "\n\n".join(part for part in extras if part) + content = result.content + if isinstance(content, str): + merged = f"{content}\n\n{suffix}" if content else suffix + else: + merged = f"{content!s}\n\n{suffix}" + return result.model_copy(update={"content": merged}) + + +def _apply_subagent_stop( + result: ToolMessage | Command[Any], + decision: SubagentStopDecision, +) -> ToolMessage | Command[Any]: + if not decision.context or not isinstance(result, ToolMessage): + return result + suffix = "\n".join(decision.context) + content = result.content + merged = ( + f"{content}\n\n{suffix}" if isinstance(content, str) and content else suffix + ) + return result.model_copy(update={"content": merged}) + + +def _inject_subagent_start_context( + request: ToolCallRequest, + decision: SubagentStartDecision, +) -> ToolCallRequest: + if not decision.context: + return request + + original = request.tool_call + raw_args = original.get("args") + args: dict[str, Any] + if isinstance(raw_args, dict): + args = {str(key): value for key, value in raw_args.items()} + else: + args = {} + description = args.get("description") + prefix = "\n".join(decision.context) + if isinstance(description, str) and description: + args["description"] = f"{prefix}\n\n{description}" + else: + args["description"] = prefix + tool_call = cast( + "ToolCall", + { + "name": str(original.get("name") or ""), + "args": args, + "id": original.get("id"), + "type": "tool_call", + }, + ) + return request.override(tool_call=tool_call) + + +def _task_agent_identity(call: ToolCallData) -> AgentIdentity: + name = call.args.get("subagent_type") + if not isinstance(name, str) or not name: + name = "unknown" + return AgentIdentity(id=call.id or name, name=name) + + +def _tool_result_text(result: ToolMessage | Command[Any]) -> str: + if isinstance(result, ToolMessage): + content = result.content + return content if isinstance(content, str) else str(content) + return "" + + +def _last_assistant_text(messages: Sequence[Any]) -> str: + for message in reversed(messages): + if isinstance(message, AIMessage): + content = message.content + if isinstance(content, str): + return content + return str(content) + return "" diff --git a/libs/code/deepagents_code/hooks/snapshot.py b/libs/code/deepagents_code/hooks/snapshot.py index d134ba66f8..ffe6105e59 100644 --- a/libs/code/deepagents_code/hooks/snapshot.py +++ b/libs/code/deepagents_code/hooks/snapshot.py @@ -173,6 +173,22 @@ def match(self, invocation: HookInvocation) -> HookMatch: ) return HookMatch(handlers=matched) + def configured_events(self) -> frozenset[HookEvent]: + """Return events that have at least one compiled handler.""" + return frozenset( + event for event, handlers in self.handlers.items() if handlers + ) + + def configured_server_events(self) -> frozenset[HookEvent]: + """Return server-owned events that have at least one compiled handler.""" + from deepagents_code.hooks.capabilities import HookOwner + + return frozenset( + event + for event in self.configured_events() + if get_event_spec(event).owner is HookOwner.SERVER + ) + def _compile_matcher( value: str | None, diff --git a/libs/code/deepagents_code/tui/textual_adapter.py b/libs/code/deepagents_code/tui/textual_adapter.py index b812a6df39..495654b9c1 100644 --- a/libs/code/deepagents_code/tui/textual_adapter.py +++ b/libs/code/deepagents_code/tui/textual_adapter.py @@ -940,6 +940,7 @@ def _notify_user_visible_output_started() -> None: suppress_resumed_output = False pending_interrupts: dict[str, tuple[tuple[Any, ...], HITLRequest]] = {} pending_ask_user: dict[str, AskUserRequest] = {} + pending_hook_resumes: dict[str, dict[str, Any]] = {} if context is None: context = CLIContext() @@ -1006,6 +1007,14 @@ def _notify_user_visible_output_started() -> None: context["approval_mode_key"] = live_key session_state.approval_mode_key = live_key + from deepagents_code.hooks.context import apply_hooks_context + + apply_hooks_context( + context, + getattr(session_state, "hooks_runtime", None), + prompt_id=getattr(session_state, "turn_id", None), + ) + # Show the Thinking spinner before each astream iteration so # both the first turn and HITL/ask_user resumes surface feedback # while the model processes input. Skip when @@ -1135,6 +1144,35 @@ def _notify_user_visible_output_started() -> None: if interrupts: for interrupt_obj in interrupts: iv = interrupt_obj.value + from deepagents_code.hooks.interrupt import ( + is_hook_interrupt_payload, + ) + + if is_hook_interrupt_payload(iv): + hooks_runtime = getattr( + session_state, "hooks_runtime", None + ) + if hooks_runtime is None: + msg = ( + "Received hook invocation interrupt " + "without a HooksRuntime" + ) + raise RuntimeError(msg) + from deepagents_code.hooks.client import ( + fulfill_hook_interrupt, + ) + + resume_value = await fulfill_hook_interrupt( + hooks_runtime, iv + ) + if resume_value is None: + msg = "Failed to parse hook interrupt" + raise RuntimeError(msg) + pending_hook_resumes[interrupt_obj.id] = ( + resume_value + ) + interrupt_occurred = True + continue if ( isinstance(iv, dict) and iv.get("type") == "ask_user" @@ -1694,7 +1732,7 @@ def _notify_user_visible_output_started() -> None: if interrupt_occurred: any_rejected = False ask_user_cancelled = False - resume_payload: dict[str, Any] = {} + resume_payload: dict[str, Any] = dict(pending_hook_resumes) # Tools mounted above start their spinner immediately, but a # tool blocked on HITL approval or `ask_user` input is not diff --git a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py new file mode 100644 index 0000000000..c57f579b0b --- /dev/null +++ b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py @@ -0,0 +1,191 @@ +"""Unit tests for Hooks v2 server-owned lifecycle integration.""" + +from __future__ import annotations + +from datetime import UTC, datetime +from pathlib import Path +from typing import Any +from uuid import uuid4 + +import pytest +from langchain_core.messages import ToolMessage + +from deepagents_code.approval_mode import ApprovalMode +from deepagents_code.hooks.client import fulfill_hook_invocation +from deepagents_code.hooks.context import apply_hooks_context +from deepagents_code.hooks.interrupt import ( + HOOK_INVOCATION_INTERRUPT_TYPE, + build_hook_interrupt_payload, + build_hook_resume_value, + is_hook_interrupt_payload, + parse_hook_interrupt_payload, + parse_hook_resume_value, +) +from deepagents_code.hooks.models.adapters import HOOKS_CONFIG_ADAPTER +from deepagents_code.hooks.models.config import HooksConfig +from deepagents_code.hooks.models.domain import ( + HookContext, + HookEvent, + HookInvocation, + PermissionEffect, + PreToolUseDecision, + PreToolUseEvent, + ToolCallData, +) +from deepagents_code.hooks.models.transport import ( + HookInvocationRequest, + HookInvocationResponse, +) +from deepagents_code.hooks.runtime import HooksRuntime +from deepagents_code.hooks.server_middleware import ( + _denied_tool_message, + _session_gate, +) +from deepagents_code.hooks.snapshot import HooksSnapshot + + +def _request(event: PreToolUseEvent | None = None) -> HookInvocationRequest: + invocation = HookInvocation( + context=HookContext( + thread_id="thread-1", + cwd=Path("/tmp"), + approval_mode=ApprovalMode.MANUAL, + ), + event=event + or PreToolUseEvent( + event=HookEvent.PRE_TOOL_USE, + call=ToolCallData(id="call-1", name="execute", args={"command": "ls"}), + ), + ) + return HookInvocationRequest( + protocol_version=1, + invocation_id=uuid4(), + snapshot_id="snapshot-1", + run_id="run-1", + invocation=invocation, + deadline=datetime(2026, 7, 23, tzinfo=UTC), + ) + + +def test_hook_interrupt_payload_round_trip() -> None: + request = _request() + payload = build_hook_interrupt_payload(request) + + assert payload["type"] == HOOK_INVOCATION_INTERRUPT_TYPE + assert is_hook_interrupt_payload(payload) + assert parse_hook_interrupt_payload(payload) == request + assert parse_hook_interrupt_payload({"type": "ask_user"}) is None + + +def test_hook_resume_value_validates_identity() -> None: + request = _request() + response = HookInvocationResponse( + protocol_version=1, + invocation_id=request.invocation_id, + snapshot_id=request.snapshot_id, + decision=PreToolUseDecision( + event=HookEvent.PRE_TOOL_USE, + permission=PermissionEffect(behavior="allow"), + ), + ) + resume = build_hook_resume_value(response) + parsed = parse_hook_resume_value( + resume, + invocation_id=request.invocation_id, + snapshot_id=request.snapshot_id, + ) + assert parsed == response + + with pytest.raises(ValueError, match="invocation_id mismatch"): + parse_hook_resume_value( + resume, + invocation_id=uuid4(), + snapshot_id=request.snapshot_id, + ) + + +def test_apply_hooks_context_sets_server_events(tmp_path: Path) -> None: + config_dir = tmp_path / "config" + config_dir.mkdir() + (config_dir / "hooks.json").write_text( + '{"hooks":{"PreToolUse":[{"hooks":[{"type":"command","command":"true"}]}]}}', + encoding="utf-8", + ) + runtime = HooksRuntime.create(cwd=tmp_path, config_dir=config_dir) + context: dict[str, Any] = {} + apply_hooks_context(context, runtime, prompt_id="prompt-1") + + assert context["hooks_snapshot_id"] == runtime.snapshot_id + assert context["hooks_server_events"] == ["PreToolUse"] + assert context["prompt_id"] == "prompt-1" + assert runtime.configured_server_events() == ("PreToolUse",) + + +def test_session_gate_requires_snapshot_and_events() -> None: + assert _session_gate(None) is None + assert _session_gate({"hooks_snapshot_id": "abc"}) is None + gate = _session_gate( + { + "hooks_snapshot_id": "abc", + "hooks_server_events": ["PreToolUse", "Stop"], + } + ) + assert gate is not None + assert gate["snapshot_id"] == "abc" + assert gate["events"] == frozenset({"PreToolUse", "Stop"}) + + +def test_denied_tool_message_for_deny_and_ask() -> None: + call = ToolCallData(id="c1", name="execute", args={}) + denied = _denied_tool_message( + call, PermissionEffect(behavior="deny", reason="nope") + ) + assert isinstance(denied, ToolMessage) + assert denied.status == "error" + assert "nope" in str(denied.content) + + asked = _denied_tool_message(call, PermissionEffect(behavior="ask")) + assert isinstance(asked, ToolMessage) + assert asked.status == "error" + + assert _denied_tool_message(call, PermissionEffect(behavior="allow")) is None + + +async def test_fulfill_hook_invocation_runs_engine(tmp_path: Path) -> None: + config_dir = tmp_path / "config" + config_dir.mkdir() + (config_dir / "hooks.json").write_text('{"hooks":{}}', encoding="utf-8") + runtime = HooksRuntime.create(cwd=tmp_path, config_dir=config_dir) + request = _request() + request = request.model_copy(update={"snapshot_id": runtime.snapshot_id}) + + resume = await fulfill_hook_invocation(runtime, request) + response = parse_hook_resume_value( + resume, + invocation_id=request.invocation_id, + snapshot_id=runtime.snapshot_id, + ) + assert isinstance(response.decision, PreToolUseDecision) + assert response.decision.permission.behavior in {"allow", "none"} + + +def test_snapshot_configured_server_events() -> None: + config = HOOKS_CONFIG_ADAPTER.validate_python( + { + "hooks": { + "SessionStart": [ + {"hooks": [{"type": "command", "command": "echo client"}]} + ], + "PreToolUse": [ + {"hooks": [{"type": "command", "command": "echo server"}]} + ], + } + } + ) + assert isinstance(config, HooksConfig) + snapshot = HooksSnapshot.from_config(config) + assert snapshot.configured_events() == { + HookEvent.SESSION_START, + HookEvent.PRE_TOOL_USE, + } + assert snapshot.configured_server_events() == {HookEvent.PRE_TOOL_USE} From 3c4a411cbcb9e2e3b60cb4b1ede8b00b52178bb8 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 23 Jul 2026 15:13:22 +0000 Subject: [PATCH 02/14] style(code): format HooksSnapshot configured_events Co-authored-by: Johannes du Plessis --- libs/code/deepagents_code/hooks/snapshot.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/libs/code/deepagents_code/hooks/snapshot.py b/libs/code/deepagents_code/hooks/snapshot.py index ffe6105e59..ead6d3aee4 100644 --- a/libs/code/deepagents_code/hooks/snapshot.py +++ b/libs/code/deepagents_code/hooks/snapshot.py @@ -175,9 +175,7 @@ def match(self, invocation: HookInvocation) -> HookMatch: def configured_events(self) -> frozenset[HookEvent]: """Return events that have at least one compiled handler.""" - return frozenset( - event for event, handlers in self.handlers.items() if handlers - ) + return frozenset(event for event, handlers in self.handlers.items() if handlers) def configured_server_events(self) -> frozenset[HookEvent]: """Return server-owned events that have at least one compiled handler.""" From 34dd24c5ed21a6d05228b21020cdb33b704474f6 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 23 Jul 2026 15:13:44 +0000 Subject: [PATCH 03/14] fix(code): satisfy hooks server-lifecycle type checks Co-authored-by: Johannes du Plessis --- libs/code/tests/unit_tests/hooks/test_server_lifecycle.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py index c57f579b0b..ba7f34fbac 100644 --- a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py +++ b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py @@ -4,12 +4,12 @@ from datetime import UTC, datetime from pathlib import Path -from typing import Any from uuid import uuid4 import pytest from langchain_core.messages import ToolMessage +from deepagents_code._cli_context import CLIContext from deepagents_code.approval_mode import ApprovalMode from deepagents_code.hooks.client import fulfill_hook_invocation from deepagents_code.hooks.context import apply_hooks_context @@ -112,7 +112,7 @@ def test_apply_hooks_context_sets_server_events(tmp_path: Path) -> None: encoding="utf-8", ) runtime = HooksRuntime.create(cwd=tmp_path, config_dir=config_dir) - context: dict[str, Any] = {} + context: CLIContext = {} apply_hooks_context(context, runtime, prompt_id="prompt-1") assert context["hooks_snapshot_id"] == runtime.snapshot_id From 8172f4059b76e36c8d61f6441db099e2d77653ab Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Thu, 23 Jul 2026 15:14:01 +0000 Subject: [PATCH 04/14] fix(code): move CLIContext import behind TYPE_CHECKING Co-authored-by: Johannes du Plessis --- libs/code/tests/unit_tests/hooks/test_server_lifecycle.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py index ba7f34fbac..cce90f9ec2 100644 --- a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py +++ b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py @@ -4,12 +4,12 @@ from datetime import UTC, datetime from pathlib import Path +from typing import TYPE_CHECKING from uuid import uuid4 import pytest from langchain_core.messages import ToolMessage -from deepagents_code._cli_context import CLIContext from deepagents_code.approval_mode import ApprovalMode from deepagents_code.hooks.client import fulfill_hook_invocation from deepagents_code.hooks.context import apply_hooks_context @@ -43,6 +43,9 @@ ) from deepagents_code.hooks.snapshot import HooksSnapshot +if TYPE_CHECKING: + from deepagents_code._cli_context import CLIContext + def _request(event: PreToolUseEvent | None = None) -> HookInvocationRequest: invocation = HookInvocation( From 54c2175e84eb40e9645d490dcc1adc3b97e79b88 Mon Sep 17 00:00:00 2001 From: Johannes du Plessis Date: Thu, 23 Jul 2026 14:25:56 -0700 Subject: [PATCH 05/14] fix(code): align server hooks lifecycle with design-doc MVP Apply PreToolUse ask/context/continue, surface notices, wire workspace trust, mount middleware on subagents, and harden Stop/SubagentStop decision handling. Co-authored-by: Cursor --- libs/code/deepagents_code/agent.py | 6 + libs/code/deepagents_code/app.py | 9 +- .../deepagents_code/client/non_interactive.py | 21 +- libs/code/deepagents_code/hooks/client.py | 48 +++ libs/code/deepagents_code/hooks/interrupt.py | 11 +- libs/code/deepagents_code/hooks/runtime.py | 3 +- .../hooks/server_middleware.py | 302 ++++++++++++------ libs/code/deepagents_code/hooks/snapshot.py | 4 +- .../deepagents_code/tui/textual_adapter.py | 16 +- .../unit_tests/hooks/test_server_lifecycle.py | 126 +++++++- 10 files changed, 411 insertions(+), 135 deletions(-) diff --git a/libs/code/deepagents_code/agent.py b/libs/code/deepagents_code/agent.py index 6c20c1c992..8186989776 100644 --- a/libs/code/deepagents_code/agent.py +++ b/libs/code/deepagents_code/agent.py @@ -2429,6 +2429,12 @@ def _subagent_cli_middleware( middleware.append(_GlmTerminalStallRecovery()) if restrictive_shell_allow_list is not None: middleware.append(ShellAllowListMiddleware(restrictive_shell_allow_list)) + # Server-owned hooks must wrap subagent tools too; otherwise Pre/Post + # ToolUse only fire on the parent graph. + from deepagents_code.hooks.server_middleware import ServerHooksMiddleware + + hooks_cwd = Path(effective_cwd) if effective_cwd is not None else Path.cwd() + middleware.append(ServerHooksMiddleware(cwd=hooks_cwd)) # Subagents share the on-disk filesystem backend and can edit the user # AGENTS.md, so they get the same managed onboarding-name block guard as # the main agent. Gated on memory because the block only exists when diff --git a/libs/code/deepagents_code/app.py b/libs/code/deepagents_code/app.py index 6b90a2fce4..b470c79225 100644 --- a/libs/code/deepagents_code/app.py +++ b/libs/code/deepagents_code/app.py @@ -2261,8 +2261,8 @@ def __init__( # Assign the backing field directly: the setter reads `self._thread_id` # to detect a thread change, and it isn't set yet. self._thread_id = thread_id or _new_thread_id() + # Optional session-scoped Hooks v2 client runtime. self.hooks_runtime = None - """Optional session-scoped Hooks v2 client runtime.""" @property def auto_approve(self) -> bool: @@ -4290,7 +4290,12 @@ def _create() -> TextualSessionState: thread_id=self._lc_thread_id, ) try: - state.hooks_runtime = HooksRuntime.create(cwd=Path(self._cwd)) + # Interactive sessions keep project hooks off until a dedicated + # workspace-trust prompt lands (design-doc security follow-up). + state.hooks_runtime = HooksRuntime.create( + cwd=Path(self._cwd), + workspace_trusted=False, + ) except Exception: logger.exception("Failed to create HooksRuntime; server hooks disabled") state.hooks_runtime = None diff --git a/libs/code/deepagents_code/client/non_interactive.py b/libs/code/deepagents_code/client/non_interactive.py index 9f05e77ab5..6e763f8a9e 100644 --- a/libs/code/deepagents_code/client/non_interactive.py +++ b/libs/code/deepagents_code/client/non_interactive.py @@ -87,6 +87,8 @@ from deepagents import FsToolName from langchain_core.runnables import RunnableConfig + from deepagents_code.hooks.runtime import HooksRuntime + logger = logging.getLogger(__name__) @@ -360,7 +362,7 @@ class StreamState: hook_response: dict[str, Any] = field(default_factory=dict) """Resume values for fulfilled Hooks v2 interrupts, keyed by interrupt id.""" - hooks_runtime: Any | None = None + hooks_runtime: HooksRuntime | None = None """Optional session-scoped HooksRuntime used to fulfill server hook interrupts.""" interrupt_occurred: bool = False @@ -957,19 +959,16 @@ async def _fulfill_pending_hook_interrupts(state: StreamState) -> None: """ if not state.pending_hook_interrupts: return - from deepagents_code.hooks.client import fulfill_hook_interrupt + from deepagents_code.hooks.client import fulfill_pending_hook_interrupts if state.hooks_runtime is None: msg = "Received hook invocation interrupt without a HooksRuntime" raise RuntimeError(msg) pending = dict(state.pending_hook_interrupts) state.pending_hook_interrupts.clear() - for interrupt_id, payload in pending.items(): - resume_value = await fulfill_hook_interrupt(state.hooks_runtime, payload) - if resume_value is None: - msg = f"Failed to parse hook interrupt {interrupt_id}" - raise RuntimeError(msg) - state.hook_response[interrupt_id] = resume_value + state.hook_response.update( + await fulfill_pending_hook_interrupts(state.hooks_runtime, pending) + ) def _process_hitl_interrupts(state: StreamState, console: Console) -> None: @@ -1149,7 +1148,11 @@ async def _run_agent_loop( from deepagents_code.hooks.runtime import HooksRuntime try: - hooks_runtime = HooksRuntime.create(cwd=Path.cwd()) + # Non-interactive mirrors Claude Code: project hooks are trusted. + hooks_runtime = HooksRuntime.create( + cwd=Path.cwd(), + workspace_trusted=True, + ) except Exception: logger.exception("Failed to create HooksRuntime; server hooks disabled") hooks_runtime = None diff --git a/libs/code/deepagents_code/hooks/client.py b/libs/code/deepagents_code/hooks/client.py index b98880cb70..dae168b435 100644 --- a/libs/code/deepagents_code/hooks/client.py +++ b/libs/code/deepagents_code/hooks/client.py @@ -2,6 +2,8 @@ from __future__ import annotations +import logging +import sys from typing import TYPE_CHECKING from deepagents_code.hooks.interrupt import ( @@ -11,9 +13,14 @@ from deepagents_code.hooks.models.transport import HookInvocationResponse if TYPE_CHECKING: + from collections.abc import Mapping + + from deepagents_code.hooks.models.domain import HookDecision from deepagents_code.hooks.models.transport import HookInvocationRequest from deepagents_code.hooks.runtime import HooksRuntime +logger = logging.getLogger(__name__) + async def fulfill_hook_invocation( runtime: HooksRuntime, @@ -39,6 +46,7 @@ async def fulfill_hook_invocation( raise ValueError(msg) decision = await runtime.invoke(request.invocation) + _apply_client_side_effects(decision) response = HookInvocationResponse( protocol_version=1, invocation_id=request.invocation_id, @@ -65,3 +73,43 @@ async def fulfill_hook_interrupt( if request is None: return None return await fulfill_hook_invocation(runtime, request) + + +async def fulfill_pending_hook_interrupts( + runtime: HooksRuntime, + pending: Mapping[str, object], +) -> dict[str, dict[str, object]]: + """Fulfill pending hook interrupts into a resume map keyed by interrupt id. + + Args: + runtime: Session-scoped client Hooks runtime. + pending: Mapping of LangGraph interrupt id to raw interrupt payload. + + Returns: + Resume values ready for `Command(resume=...)`. + + Raises: + RuntimeError: If a payload is not a valid hook interrupt. + """ + resumes: dict[str, dict[str, object]] = {} + for interrupt_id, payload in pending.items(): + resume_value = await fulfill_hook_interrupt(runtime, payload) + if resume_value is None: + msg = f"Failed to parse hook interrupt {interrupt_id}" + raise RuntimeError(msg) + resumes[interrupt_id] = resume_value + return resumes + + +def _apply_client_side_effects(decision: HookDecision) -> None: + """Surface user notices and emit validated terminal sequences. + + `systemMessage` must never become model context; notices are logged for the + operator. Terminal sequences were allowlisted in the reducer. + """ + for notice in decision.user_notices: + logger.warning("Hook user notice: %s", notice) + for sequence in decision.terminal_sequences: + sys.stdout.write(sequence) + if decision.terminal_sequences: + sys.stdout.flush() diff --git a/libs/code/deepagents_code/hooks/interrupt.py b/libs/code/deepagents_code/hooks/interrupt.py index aaba8e94ac..6e964e62a4 100644 --- a/libs/code/deepagents_code/hooks/interrupt.py +++ b/libs/code/deepagents_code/hooks/interrupt.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Any, Literal, TypeAlias -from pydantic import BaseModel, ConfigDict, Field, TypeAdapter +from pydantic import BaseModel, ConfigDict, TypeAdapter from deepagents_code.hooks.models.adapters import ( HOOK_INVOCATION_RESPONSE_ADAPTER, @@ -119,12 +119,3 @@ def is_hook_interrupt_payload(value: object) -> bool: return ( isinstance(value, dict) and value.get("type") == HOOK_INVOCATION_INTERRUPT_TYPE ) - - -class HookInterruptMismatch(BaseModel): - """Diagnostic retained when a hook resume cannot be applied.""" - - model_config = ConfigDict(extra="forbid") - - code: Literal["hook_resume_mismatch"] = "hook_resume_mismatch" - message: str = Field(min_length=1) diff --git a/libs/code/deepagents_code/hooks/runtime.py b/libs/code/deepagents_code/hooks/runtime.py index c383866191..351c7a003c 100644 --- a/libs/code/deepagents_code/hooks/runtime.py +++ b/libs/code/deepagents_code/hooks/runtime.py @@ -42,7 +42,8 @@ class HooksRuntime: """Client-owned session runtime around an immutable Hooks snapshot. Owns configuration snapshot identity, transcript materialization, and the - `HookEngine`. Lifecycle call sites are intentionally not wired here. + `HookEngine`. Server-owned lifecycle events reach this runtime through the + interrupt fulfill path in `hooks.client`. """ snapshot: HooksSnapshot diff --git a/libs/code/deepagents_code/hooks/server_middleware.py b/libs/code/deepagents_code/hooks/server_middleware.py index e56c7423ee..2dfb8282bb 100644 --- a/libs/code/deepagents_code/hooks/server_middleware.py +++ b/libs/code/deepagents_code/hooks/server_middleware.py @@ -9,10 +9,16 @@ import time from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field from datetime import UTC, datetime, timedelta -from typing import TYPE_CHECKING, Any, NotRequired, cast +from typing import TYPE_CHECKING, Any, NotRequired, TypeVar, cast from uuid import UUID, uuid4 +from langchain.agents.middleware.human_in_the_loop import ( + ActionRequest, + HITLRequest, + ReviewConfig, +) from langchain.agents.middleware.types import ( AgentMiddleware, AgentState, @@ -31,6 +37,7 @@ ) from deepagents_code.hooks.models.domain import ( AgentIdentity, + BaseHookDecision, HookContext, HookDecision, HookEvent, @@ -77,6 +84,14 @@ class _SessionHookGate(TypedDict): events: frozenset[str] +@dataclass(slots=True) +class _PreToolOutcome: + """PreToolUse gate result for the tool-call wrapper.""" + + blocked: ToolMessage | None = None + context: tuple[str, ...] = field(default_factory=tuple) + + class ServerHooksMiddleware(AgentMiddleware[ServerHooksState, ContextT, ResponseT]): """Emit server-owned lifecycle events over the hook interrupt transport.""" @@ -108,7 +123,25 @@ def wrap_tool_call( Returns: Tool result, possibly rewritten by hook decisions. """ - return self._run_tool_call(request, handler, async_handler=False) + gate = _session_gate(request.runtime.context) + call = _tool_call_data(request) + context = _hook_context( + request.runtime.context, request.runtime.config, self._cwd + ) + request = self._maybe_subagent_start(request, call, context, gate) + pre = self._maybe_pre_tool_use(call, context, gate, request.runtime.config) + if pre.blocked is not None: + return _append_message_text(pre.blocked, pre.context) + started = time.perf_counter() + result = handler(request) + duration_ms = int((time.perf_counter() - started) * 1000) + result = _append_message_text(result, pre.context) + result = self._maybe_post_tool_use( + call, context, gate, request.runtime.config, result, duration_ms + ) + return self._maybe_subagent_stop( + call, context, gate, request.runtime.config, result + ) async def awrap_tool_call( self, @@ -120,7 +153,25 @@ async def awrap_tool_call( Returns: Tool result, possibly rewritten by hook decisions. """ - return await self._run_tool_call_async(request, handler) + gate = _session_gate(request.runtime.context) + call = _tool_call_data(request) + context = _hook_context( + request.runtime.context, request.runtime.config, self._cwd + ) + request = self._maybe_subagent_start(request, call, context, gate) + pre = self._maybe_pre_tool_use(call, context, gate, request.runtime.config) + if pre.blocked is not None: + return _append_message_text(pre.blocked, pre.context) + started = time.perf_counter() + result = await handler(request) + duration_ms = int((time.perf_counter() - started) * 1000) + result = _append_message_text(result, pre.context) + result = self._maybe_post_tool_use( + call, context, gate, request.runtime.config, result, duration_ms + ) + return self._maybe_subagent_stop( + call, context, gate, request.runtime.config, result + ) @hook_config(can_jump_to=["model"]) def after_agent( @@ -148,57 +199,6 @@ async def aafter_agent( """ return self._after_agent(state, runtime) - def _run_tool_call( - self, - request: ToolCallRequest, - handler: Callable[[ToolCallRequest], ToolMessage | Command[Any]], - *, - async_handler: bool, - ) -> ToolMessage | Command[Any]: - del async_handler - gate = _session_gate(request.runtime.context) - call = _tool_call_data(request) - context = _hook_context( - request.runtime.context, request.runtime.config, self._cwd - ) - request = self._maybe_subagent_start(request, call, context, gate) - denied = self._maybe_pre_tool_use(call, context, gate, request.runtime.config) - if denied is not None: - return denied - started = time.perf_counter() - result = handler(request) - duration_ms = int((time.perf_counter() - started) * 1000) - result = self._maybe_post_tool_use( - call, context, gate, request.runtime.config, result, duration_ms - ) - return self._maybe_subagent_stop( - call, context, gate, request.runtime.config, result - ) - - async def _run_tool_call_async( - self, - request: ToolCallRequest, - handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command[Any]]], - ) -> ToolMessage | Command[Any]: - gate = _session_gate(request.runtime.context) - call = _tool_call_data(request) - context = _hook_context( - request.runtime.context, request.runtime.config, self._cwd - ) - request = self._maybe_subagent_start(request, call, context, gate) - denied = self._maybe_pre_tool_use(call, context, gate, request.runtime.config) - if denied is not None: - return denied - started = time.perf_counter() - result = await handler(request) - duration_ms = int((time.perf_counter() - started) * 1000) - result = self._maybe_post_tool_use( - call, context, gate, request.runtime.config, result, duration_ms - ) - return self._maybe_subagent_stop( - call, context, gate, request.runtime.config, result - ) - def _maybe_subagent_start( self, request: ToolCallRequest, @@ -218,9 +218,21 @@ def _maybe_subagent_start( config=request.runtime.config, deadline=self._default_deadline, ) - if not isinstance(decision, SubagentStartDecision): - msg = "Expected SubagentStartDecision, got {type(decision).__name__}" - raise TypeError(msg) + decision = _require_decision(decision, SubagentStartDecision) + if not decision.continue_processing: + # SubagentStart has no deny ToolMessage path; refuse spawn by + # clearing the description so the task tool fails closed upstream. + return _inject_subagent_start_context( + request, + SubagentStartDecision( + event=HookEvent.SUBAGENT_START, + context=[ + decision.stop_reason or "Blocked by SubagentStart hook", + *decision.context, + ], + continue_processing=False, + ), + ) return _inject_subagent_start_context(request, decision) def _maybe_pre_tool_use( @@ -229,9 +241,9 @@ def _maybe_pre_tool_use( context: HookContext, gate: _SessionHookGate | None, config: Mapping[str, Any] | None, - ) -> ToolMessage | None: + ) -> _PreToolOutcome: if not _event_enabled(gate, HookEvent.PRE_TOOL_USE): - return None + return _PreToolOutcome() decision = _invoke_hook( context, PreToolUseEvent(event=HookEvent.PRE_TOOL_USE, call=call), @@ -239,10 +251,29 @@ def _maybe_pre_tool_use( config=config, deadline=self._default_deadline, ) - if not isinstance(decision, PreToolUseDecision): - msg = "Expected PreToolUseDecision, got {type(decision).__name__}" - raise TypeError(msg) - return _denied_tool_message(call, decision.permission) + decision = _require_decision(decision, PreToolUseDecision) + context_parts = tuple(decision.context) + if not decision.continue_processing: + return _PreToolOutcome( + blocked=_denied_tool_message( + call, + PermissionEffect( + behavior="deny", + reason=decision.stop_reason or "Stopped by PreToolUse hook", + ), + ), + context=context_parts, + ) + behavior = decision.permission.behavior + if behavior == "deny": + return _PreToolOutcome( + blocked=_denied_tool_message(call, decision.permission), + context=context_parts, + ) + if behavior == "ask": + blocked = _ask_permission_via_hitl(call, decision.permission) + return _PreToolOutcome(blocked=blocked, context=context_parts) + return _PreToolOutcome(context=context_parts) def _maybe_post_tool_use( self, @@ -269,9 +300,7 @@ def _maybe_post_tool_use( config=config, deadline=self._default_deadline, ) - if not isinstance(decision, PostToolUseDecision): - msg = "Expected PostToolUseDecision, got {type(decision).__name__}" - raise TypeError(msg) + decision = _require_decision(decision, PostToolUseDecision) return _apply_post_tool_use(result, decision) def _maybe_subagent_stop( @@ -299,9 +328,7 @@ def _maybe_subagent_stop( config=config, deadline=self._default_deadline, ) - if not isinstance(decision, SubagentStopDecision): - msg = "Expected SubagentStopDecision, got {type(decision).__name__}" - raise TypeError(msg) + decision = _require_decision(decision, SubagentStopDecision) return _apply_subagent_stop(result, decision) def _after_agent( @@ -325,10 +352,11 @@ def _after_agent( config=None, deadline=self._default_deadline, ) - if not isinstance(decision, StopDecision): - msg = "Expected StopDecision, got {type(decision).__name__}" - raise TypeError(msg) - if not decision.continue_loop: + decision = _require_decision(decision, StopDecision) + if not decision.continue_processing or not decision.continue_loop: + # Reset so a later independent turn does not inherit the count. + if continuation: + return {_STOP_STATE_KEY: 0} return None feedback = "\n".join(decision.feedback).strip() or ( decision.stop_reason or "Continue working." @@ -340,6 +368,19 @@ def _after_agent( } +_DecisionT = TypeVar("_DecisionT", bound=BaseHookDecision) + + +def _require_decision( + decision: HookDecision, + expected: type[_DecisionT], +) -> _DecisionT: + if not isinstance(decision, expected): + msg = f"Expected {expected.__name__}, got {type(decision).__name__}" + raise TypeError(msg) + return decision + + def _session_gate(runtime_context: object) -> _SessionHookGate | None: fields = _context_mapping(runtime_context) snapshot_id = fields.get("hooks_snapshot_id") @@ -415,6 +456,14 @@ def _hook_context( def _context_mapping(runtime_context: object) -> dict[str, Any]: + """Project LangGraph run context (dataclass or mapping) into a plain dict. + + In-process graphs coerce `context=` into `CLIContextSchema`; RemoteGraph + delivers a plain mapping. Both shapes are accepted here. + + Returns: + A shallow string-keyed dict of the hook-relevant context fields. + """ if runtime_context is None: return {} if isinstance(runtime_context, Mapping): @@ -486,14 +535,8 @@ def _mcp_server_from_tool(tool: object | None) -> str | None: def _denied_tool_message( call: ToolCallData, permission: PermissionEffect, -) -> ToolMessage | None: - if permission.behavior not in {"deny", "ask"}: - return None - reason = permission.reason or ( - "Blocked by PreToolUse hook" - if permission.behavior == "deny" - else "PreToolUse hook requested approval (ask is not applied yet)" - ) +) -> ToolMessage: + reason = permission.reason or "Blocked by PreToolUse hook" wire_name = to_wire_tool_name(call.name, mcp_server=call.mcp_server) return ToolMessage( content=f"{wire_name} blocked by hook: {reason}", @@ -503,6 +546,74 @@ def _denied_tool_message( ) +def _ask_permission_via_hitl( + call: ToolCallData, + permission: PermissionEffect, +) -> ToolMessage | None: + """Escalate PreToolUse `ask` through the existing HITL interrupt channel. + + Returns: + A deny ToolMessage when the user rejects, otherwise `None` to proceed. + """ + description = permission.reason or "PreToolUse hook requested approval" + response = interrupt( + HITLRequest( + action_requests=[ + ActionRequest( + name=call.name, + args=dict(call.args), + description=description, + ) + ], + review_configs=[ + ReviewConfig( + action_name=call.name, + allowed_decisions=["approve", "reject"], + ) + ], + ) + ) + decisions: Sequence[Any] + if isinstance(response, Mapping): + raw = response.get("decisions", ()) + decisions = raw if isinstance(raw, Sequence) else () + else: + decisions = () + if not decisions: + return _denied_tool_message( + call, + PermissionEffect( + behavior="deny", + reason="PreToolUse ask was not answered", + ), + ) + first = decisions[0] + decision_type = first.get("type") if isinstance(first, Mapping) else None + if decision_type != "approve": + reject_message = None + if isinstance(first, Mapping): + raw_message = first.get("message") + if isinstance(raw_message, str) and raw_message: + reject_message = raw_message + return _denied_tool_message( + call, + PermissionEffect( + behavior="deny", + reason=reject_message or description, + ), + ) + return None + + +def _append_message_text( + result: ToolMessage | Command[Any], + parts: Sequence[str], +) -> ToolMessage | Command[Any]: + if not parts or not isinstance(result, ToolMessage): + return result + return _merge_tool_message_content(result, "\n".join(parts)) + + def _apply_post_tool_use( result: ToolMessage, decision: PostToolUseDecision, @@ -512,15 +623,13 @@ def _apply_post_tool_use( extras.append("\n".join(decision.feedback)) if decision.context: extras.append("\n".join(decision.context)) + if decision.stop_reason and not decision.continue_processing: + extras.append(decision.stop_reason) if not extras: return result - suffix = "\n\n".join(part for part in extras if part) - content = result.content - if isinstance(content, str): - merged = f"{content}\n\n{suffix}" if content else suffix - else: - merged = f"{content!s}\n\n{suffix}" - return result.model_copy(update={"content": merged}) + return _merge_tool_message_content( + result, "\n\n".join(part for part in extras if part) + ) def _apply_subagent_stop( @@ -529,11 +638,20 @@ def _apply_subagent_stop( ) -> ToolMessage | Command[Any]: if not decision.context or not isinstance(result, ToolMessage): return result - suffix = "\n".join(decision.context) + return _merge_tool_message_content(result, "\n".join(decision.context)) + + +def _merge_tool_message_content(result: ToolMessage, suffix: str) -> ToolMessage: + if not suffix: + return result content = result.content - merged = ( - f"{content}\n\n{suffix}" if isinstance(content, str) and content else suffix - ) + if isinstance(content, str): + merged = f"{content}\n\n{suffix}" if content else suffix + # Preserve structured content blocks; append a text block. + elif isinstance(content, list): + merged = [*content, {"type": "text", "text": suffix}] + else: + merged = f"{content!s}\n\n{suffix}" return result.model_copy(update={"content": merged}) diff --git a/libs/code/deepagents_code/hooks/snapshot.py b/libs/code/deepagents_code/hooks/snapshot.py index ead6d3aee4..fa6bdeb9b3 100644 --- a/libs/code/deepagents_code/hooks/snapshot.py +++ b/libs/code/deepagents_code/hooks/snapshot.py @@ -8,7 +8,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING -from deepagents_code.hooks.capabilities import get_event_spec +from deepagents_code.hooks.capabilities import HookOwner, get_event_spec from deepagents_code.hooks.loading import compute_snapshot_id from deepagents_code.hooks.models.domain import ( HookDiagnostic, @@ -179,8 +179,6 @@ def configured_events(self) -> frozenset[HookEvent]: def configured_server_events(self) -> frozenset[HookEvent]: """Return server-owned events that have at least one compiled handler.""" - from deepagents_code.hooks.capabilities import HookOwner - return frozenset( event for event in self.configured_events() diff --git a/libs/code/deepagents_code/tui/textual_adapter.py b/libs/code/deepagents_code/tui/textual_adapter.py index 495654b9c1..17e5f6aee2 100644 --- a/libs/code/deepagents_code/tui/textual_adapter.py +++ b/libs/code/deepagents_code/tui/textual_adapter.py @@ -1007,11 +1007,13 @@ def _notify_user_visible_output_started() -> None: context["approval_mode_key"] = live_key session_state.approval_mode_key = live_key + from deepagents_code.hooks.client import fulfill_hook_interrupt from deepagents_code.hooks.context import apply_hooks_context + from deepagents_code.hooks.interrupt import is_hook_interrupt_payload apply_hooks_context( context, - getattr(session_state, "hooks_runtime", None), + session_state.hooks_runtime, prompt_id=getattr(session_state, "turn_id", None), ) @@ -1144,24 +1146,14 @@ def _notify_user_visible_output_started() -> None: if interrupts: for interrupt_obj in interrupts: iv = interrupt_obj.value - from deepagents_code.hooks.interrupt import ( - is_hook_interrupt_payload, - ) - if is_hook_interrupt_payload(iv): - hooks_runtime = getattr( - session_state, "hooks_runtime", None - ) + hooks_runtime = session_state.hooks_runtime if hooks_runtime is None: msg = ( "Received hook invocation interrupt " "without a HooksRuntime" ) raise RuntimeError(msg) - from deepagents_code.hooks.client import ( - fulfill_hook_interrupt, - ) - resume_value = await fulfill_hook_interrupt( hooks_runtime, iv ) diff --git a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py index cce90f9ec2..31f564a6e4 100644 --- a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py +++ b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py @@ -4,7 +4,8 @@ from datetime import UTC, datetime from pathlib import Path -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any +from unittest.mock import MagicMock from uuid import uuid4 import pytest @@ -28,8 +29,11 @@ HookEvent, HookInvocation, PermissionEffect, + PostToolUseDecision, PreToolUseDecision, PreToolUseEvent, + StopDecision, + SubagentStopDecision, ToolCallData, ) from deepagents_code.hooks.models.transport import ( @@ -38,7 +42,13 @@ ) from deepagents_code.hooks.runtime import HooksRuntime from deepagents_code.hooks.server_middleware import ( + ServerHooksMiddleware, + _append_message_text, + _apply_post_tool_use, + _apply_subagent_stop, + _ask_permission_via_hitl, _denied_tool_message, + _merge_tool_message_content, _session_gate, ) from deepagents_code.hooks.snapshot import HooksSnapshot @@ -138,7 +148,7 @@ def test_session_gate_requires_snapshot_and_events() -> None: assert gate["events"] == frozenset({"PreToolUse", "Stop"}) -def test_denied_tool_message_for_deny_and_ask() -> None: +def test_denied_tool_message_for_deny() -> None: call = ToolCallData(id="c1", name="execute", args={}) denied = _denied_tool_message( call, PermissionEffect(behavior="deny", reason="nope") @@ -147,11 +157,115 @@ def test_denied_tool_message_for_deny_and_ask() -> None: assert denied.status == "error" assert "nope" in str(denied.content) - asked = _denied_tool_message(call, PermissionEffect(behavior="ask")) - assert isinstance(asked, ToolMessage) - assert asked.status == "error" - assert _denied_tool_message(call, PermissionEffect(behavior="allow")) is None +def test_merge_tool_message_preserves_structured_content() -> None: + result = ToolMessage( + content=[{"type": "text", "text": "parent result"}], + tool_call_id="c1", + name="task", + ) + merged = _merge_tool_message_content(result, "hook context") + assert isinstance(merged.content, list) + assert merged.content[0] == {"type": "text", "text": "parent result"} + assert merged.content[-1] == {"type": "text", "text": "hook context"} + + +def test_apply_subagent_stop_preserves_structured_content() -> None: + result = ToolMessage( + content=[{"type": "text", "text": "done"}], + tool_call_id="c1", + name="task", + ) + updated = _apply_subagent_stop( + result, + SubagentStopDecision( + event=HookEvent.SUBAGENT_STOP, + context=["extra"], + ), + ) + assert isinstance(updated, ToolMessage) + assert isinstance(updated.content, list) + assert "extra" in str(updated.content[-1]) + + +def test_apply_post_tool_use_appends_feedback_and_context() -> None: + result = ToolMessage(content="ok", tool_call_id="c1", name="execute") + updated = _apply_post_tool_use( + result, + PostToolUseDecision( + event=HookEvent.POST_TOOL_USE, + feedback=["fix it"], + context=["note"], + ), + ) + assert "ok" in str(updated.content) + assert "fix it" in str(updated.content) + assert "note" in str(updated.content) + + +def test_append_pretool_context_to_result() -> None: + result = ToolMessage(content="ran", tool_call_id="c1", name="execute") + updated = _append_message_text(result, ("pre context",)) + assert isinstance(updated, ToolMessage) + assert "ran" in str(updated.content) + assert "pre context" in str(updated.content) + + +def test_ask_permission_via_hitl_approve(monkeypatch: pytest.MonkeyPatch) -> None: + call = ToolCallData(id="c1", name="execute", args={"command": "ls"}) + + def _fake_interrupt(payload: object) -> dict[str, object]: + assert isinstance(payload, dict) + return {"decisions": [{"type": "approve"}]} + + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware.interrupt", + _fake_interrupt, + ) + assert ( + _ask_permission_via_hitl(call, PermissionEffect(behavior="ask", reason="sure?")) + is None + ) + + +def test_ask_permission_via_hitl_reject(monkeypatch: pytest.MonkeyPatch) -> None: + call = ToolCallData(id="c1", name="execute", args={}) + + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware.interrupt", + lambda _payload: {"decisions": [{"type": "reject", "message": "no"}]}, + ) + blocked = _ask_permission_via_hitl(call, PermissionEffect(behavior="ask")) + assert isinstance(blocked, ToolMessage) + assert blocked.status == "error" + assert "no" in str(blocked.content) + + +def test_stop_resets_continuation_count_when_finished( + monkeypatch: pytest.MonkeyPatch, +) -> None: + middleware = ServerHooksMiddleware(cwd=Path("/tmp")) + state: dict[str, Any] = { + "messages": [], + "_hooks_stop_continuation_count": 3, + } + runtime = MagicMock() + runtime.context = { + "hooks_snapshot_id": "snap", + "hooks_server_events": ["Stop"], + "thread_id": "t1", + "approval_mode": "manual", + } + + def _fake_invoke(*_args: object, **_kwargs: object) -> StopDecision: + return StopDecision(event=HookEvent.STOP, continue_loop=False) + + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware._invoke_hook", + _fake_invoke, + ) + update = middleware._after_agent(state, runtime) + assert update == {"_hooks_stop_continuation_count": 0} async def test_fulfill_hook_invocation_runs_engine(tmp_path: Path) -> None: From 46dc9ba83e362bdc3c4f637a8f90888bc78d846d Mon Sep 17 00:00:00 2001 From: Johannes du Plessis Date: Fri, 24 Jul 2026 08:26:17 -0700 Subject: [PATCH 06/14] fix(code): harden server hooks lifecycle review findings Enforce SubagentStart denies, skip Stop on subagent graphs, and require `--trust-project-hooks` before loading project hook commands headlessly. Co-authored-by: Cursor --- libs/code/deepagents_code/agent.py | 6 +- .../deepagents_code/client/non_interactive.py | 17 ++++- .../hooks/server_middleware.py | 35 ++++++---- libs/code/deepagents_code/main.py | 9 +++ .../deepagents_code/tui/textual_adapter.py | 6 +- libs/code/deepagents_code/ui.py | 3 + .../unit_tests/hooks/test_server_lifecycle.py | 65 ++++++++++++++++- .../tests/unit_tests/test_non_interactive.py | 69 +++++++++++++++++++ 8 files changed, 188 insertions(+), 22 deletions(-) diff --git a/libs/code/deepagents_code/agent.py b/libs/code/deepagents_code/agent.py index 8186989776..8ecdc186ab 100644 --- a/libs/code/deepagents_code/agent.py +++ b/libs/code/deepagents_code/agent.py @@ -2430,11 +2430,13 @@ def _subagent_cli_middleware( if restrictive_shell_allow_list is not None: middleware.append(ShellAllowListMiddleware(restrictive_shell_allow_list)) # Server-owned hooks must wrap subagent tools too; otherwise Pre/Post - # ToolUse only fire on the parent graph. + # ToolUse only fire on the parent graph. Disable Stop so finishing a + # subagent does not emit the main-agent Stop event (SubagentStop still + # fires from the parent wrap around `task`). from deepagents_code.hooks.server_middleware import ServerHooksMiddleware hooks_cwd = Path(effective_cwd) if effective_cwd is not None else Path.cwd() - middleware.append(ServerHooksMiddleware(cwd=hooks_cwd)) + middleware.append(ServerHooksMiddleware(cwd=hooks_cwd, emit_stop=False)) # Subagents share the on-disk filesystem backend and can edit the user # AGENTS.md, so they get the same managed onboarding-name block guard as # the main agent. Gated on memory because the block only exists when diff --git a/libs/code/deepagents_code/client/non_interactive.py b/libs/code/deepagents_code/client/non_interactive.py index 6e763f8a9e..c538d660fe 100644 --- a/libs/code/deepagents_code/client/non_interactive.py +++ b/libs/code/deepagents_code/client/non_interactive.py @@ -1083,6 +1083,7 @@ async def _run_agent_loop( max_turns: int | None = None, rubric: str | None = None, show_rubric_iterations: bool = False, + trust_project_hooks: bool = False, ) -> None: """Run the agent and handle HITL interrupts until the task completes. @@ -1115,6 +1116,11 @@ async def _run_agent_loop( `None` leaves it unset (no grading). show_rubric_iterations: Whether rubric lifecycle messages should include iteration numbers. + trust_project_hooks: When `True`, load project-scoped + `.deepagents/hooks.json` command handlers. + + Defaults to `False` so untrusted checkouts cannot execute repository + hooks in CI without an explicit opt-in. Raises: HITLIterationLimitError: If the effective turn limit is exceeded. @@ -1148,10 +1154,10 @@ async def _run_agent_loop( from deepagents_code.hooks.runtime import HooksRuntime try: - # Non-interactive mirrors Claude Code: project hooks are trusted. + # Project hooks require an explicit opt-in, matching `--trust-project-mcp`. hooks_runtime = HooksRuntime.create( cwd=Path.cwd(), - workspace_trusted=True, + workspace_trusted=trust_project_hooks, ) except Exception: logger.exception("Failed to create HooksRuntime; server hooks disabled") @@ -1425,6 +1431,7 @@ async def run_non_interactive( rubric_model: str | None = None, rubric_max_iterations: int | None = None, recursion_limit: int | None = None, + trust_project_hooks: bool = False, ) -> int: """Run a single task non-interactively and exit. @@ -1501,6 +1508,11 @@ async def run_non_interactive( uses the middleware default. recursion_limit: Explicit main-agent `recursion_limit`; `None` resolves from env / `config.toml` / default at agent-build time. + trust_project_hooks: When `True`, load project-scoped + `.deepagents/hooks.json` handlers. + + Defaults to `False` so untrusted repositories cannot execute hook + commands without an explicit `--trust-project-hooks` opt-in. Returns: Exit code: 0 for success, 1 for error, 124 when the `--max-turns` @@ -1751,6 +1763,7 @@ def discover_all_skills() -> tuple[list[ExtendedSkillMetadata], list[Path]]: max_turns=max_turns, rubric=rubric, show_rubric_iterations=rubric_max_iterations is not None, + trust_project_hooks=trust_project_hooks, ) except KeyboardInterrupt: diff --git a/libs/code/deepagents_code/hooks/server_middleware.py b/libs/code/deepagents_code/hooks/server_middleware.py index 2dfb8282bb..32d07996e4 100644 --- a/libs/code/deepagents_code/hooks/server_middleware.py +++ b/libs/code/deepagents_code/hooks/server_middleware.py @@ -102,16 +102,21 @@ def __init__( *, cwd: Path, default_deadline: timedelta = _DEFAULT_DEADLINE, + emit_stop: bool = True, ) -> None: """Initialize middleware. Args: cwd: Session working directory projected into hook context. default_deadline: Client execution deadline attached to requests. + emit_stop: Whether to emit the main-agent `Stop` event from + `after_agent`. Subagent graphs set this to `False` so they still + wrap tools without firing parent `Stop` handlers. """ super().__init__() self._cwd = cwd self._default_deadline = default_deadline + self._emit_stop = emit_stop def wrap_tool_call( self, @@ -128,7 +133,10 @@ def wrap_tool_call( context = _hook_context( request.runtime.context, request.runtime.config, self._cwd ) - request = self._maybe_subagent_start(request, call, context, gate) + started_or_blocked = self._maybe_subagent_start(request, call, context, gate) + if isinstance(started_or_blocked, ToolMessage): + return started_or_blocked + request = started_or_blocked pre = self._maybe_pre_tool_use(call, context, gate, request.runtime.config) if pre.blocked is not None: return _append_message_text(pre.blocked, pre.context) @@ -158,7 +166,10 @@ async def awrap_tool_call( context = _hook_context( request.runtime.context, request.runtime.config, self._cwd ) - request = self._maybe_subagent_start(request, call, context, gate) + started_or_blocked = self._maybe_subagent_start(request, call, context, gate) + if isinstance(started_or_blocked, ToolMessage): + return started_or_blocked + request = started_or_blocked pre = self._maybe_pre_tool_use(call, context, gate, request.runtime.config) if pre.blocked is not None: return _append_message_text(pre.blocked, pre.context) @@ -205,7 +216,7 @@ def _maybe_subagent_start( call: ToolCallData, context: HookContext, gate: _SessionHookGate | None, - ) -> ToolCallRequest: + ) -> ToolCallRequest | ToolMessage: if call.name != _TASK_TOOL_NAME or not _event_enabled( gate, HookEvent.SUBAGENT_START ): @@ -220,17 +231,11 @@ def _maybe_subagent_start( ) decision = _require_decision(decision, SubagentStartDecision) if not decision.continue_processing: - # SubagentStart has no deny ToolMessage path; refuse spawn by - # clearing the description so the task tool fails closed upstream. - return _inject_subagent_start_context( - request, - SubagentStartDecision( - event=HookEvent.SUBAGENT_START, - context=[ - decision.stop_reason or "Blocked by SubagentStart hook", - *decision.context, - ], - continue_processing=False, + return _denied_tool_message( + call, + PermissionEffect( + behavior="deny", + reason=decision.stop_reason or "Blocked by SubagentStart hook", ), ) return _inject_subagent_start_context(request, decision) @@ -336,6 +341,8 @@ def _after_agent( state: ServerHooksState, runtime: Runtime[ContextT], ) -> dict[str, Any] | None: + if not self._emit_stop: + return None gate = _session_gate(runtime.context) if not _event_enabled(gate, HookEvent.STOP): return None diff --git a/libs/code/deepagents_code/main.py b/libs/code/deepagents_code/main.py index 05e5732cfb..d55b1c6a47 100644 --- a/libs/code/deepagents_code/main.py +++ b/libs/code/deepagents_code/main.py @@ -2082,6 +2082,12 @@ def help_parent(help_fn: Callable[[], None]) -> list[argparse.ArgumentParser]: help="Trust project-level MCP configs with stdio and remote servers " "(skip interactive approval prompt)", ) + parser.add_argument( + "--trust-project-hooks", + action="store_true", + help="Trust project-level `.deepagents/hooks.json` command handlers " + "(required for headless/CI runs that should load repository hooks)", + ) parser.add_argument( "--interpreter", action=argparse.BooleanOptionalAction, @@ -4551,6 +4557,9 @@ def cli_main() -> None: mcp_config_path=getattr(args, "mcp_config", None), no_mcp=getattr(args, "no_mcp", False), trust_project_mcp=getattr(args, "trust_project_mcp", False), + trust_project_hooks=getattr( + args, "trust_project_hooks", False + ), enable_interpreter=enable_interpreter, interpreter_ptc=interpreter_ptc, allow_fs_tools=allow_fs_tools, diff --git a/libs/code/deepagents_code/tui/textual_adapter.py b/libs/code/deepagents_code/tui/textual_adapter.py index 17e5f6aee2..abda43282f 100644 --- a/libs/code/deepagents_code/tui/textual_adapter.py +++ b/libs/code/deepagents_code/tui/textual_adapter.py @@ -1013,7 +1013,7 @@ def _notify_user_visible_output_started() -> None: apply_hooks_context( context, - session_state.hooks_runtime, + getattr(session_state, "hooks_runtime", None), prompt_id=getattr(session_state, "turn_id", None), ) @@ -1147,7 +1147,9 @@ def _notify_user_visible_output_started() -> None: for interrupt_obj in interrupts: iv = interrupt_obj.value if is_hook_interrupt_payload(iv): - hooks_runtime = session_state.hooks_runtime + hooks_runtime = getattr( + session_state, "hooks_runtime", None + ) if hooks_runtime is None: msg = ( "Received hook invocation interrupt " diff --git a/libs/code/deepagents_code/ui.py b/libs/code/deepagents_code/ui.py index 4ff1c37b24..711b03b330 100644 --- a/libs/code/deepagents_code/ui.py +++ b/libs/code/deepagents_code/ui.py @@ -171,6 +171,9 @@ def show_help() -> None: console.print( " --trust-project-mcp Trust project MCP configs (skip approval prompt)" ) + console.print( + " --trust-project-hooks Trust project hooks.json command handlers" + ) console.print( " --interpreter, --no-interpreter" " Enable or disable JS interpreter (`js_eval`) middleware" diff --git a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py index 31f564a6e4..bc2f5da6ea 100644 --- a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py +++ b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py @@ -4,7 +4,7 @@ from datetime import UTC, datetime from pathlib import Path -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING from unittest.mock import MagicMock from uuid import uuid4 @@ -43,6 +43,7 @@ from deepagents_code.hooks.runtime import HooksRuntime from deepagents_code.hooks.server_middleware import ( ServerHooksMiddleware, + ServerHooksState, _append_message_text, _apply_post_tool_use, _apply_subagent_stop, @@ -245,7 +246,7 @@ def test_stop_resets_continuation_count_when_finished( monkeypatch: pytest.MonkeyPatch, ) -> None: middleware = ServerHooksMiddleware(cwd=Path("/tmp")) - state: dict[str, Any] = { + state: ServerHooksState = { "messages": [], "_hooks_stop_continuation_count": 3, } @@ -268,6 +269,66 @@ def _fake_invoke(*_args: object, **_kwargs: object) -> StopDecision: assert update == {"_hooks_stop_continuation_count": 0} +def test_emit_stop_false_skips_after_agent( + monkeypatch: pytest.MonkeyPatch, +) -> None: + middleware = ServerHooksMiddleware(cwd=Path("/tmp"), emit_stop=False) + state: ServerHooksState = {"messages": []} + runtime = MagicMock() + runtime.context = { + "hooks_snapshot_id": "snap", + "hooks_server_events": ["Stop"], + "thread_id": "t1", + "approval_mode": "manual", + } + invoke = MagicMock() + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware._invoke_hook", + invoke, + ) + assert middleware._after_agent(state, runtime) is None + invoke.assert_not_called() + + +def test_subagent_start_deny_returns_error_tool_message( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from deepagents_code.hooks.models.domain import SubagentStartDecision + + middleware = ServerHooksMiddleware(cwd=Path("/tmp")) + request = MagicMock() + request.tool_call = { + "name": "task", + "args": {"subagent_type": "researcher", "description": "go"}, + "id": "call-1", + "type": "tool_call", + } + request.tool = None + request.runtime.context = { + "hooks_snapshot_id": "snap", + "hooks_server_events": ["SubagentStart"], + "thread_id": "t1", + "approval_mode": "manual", + } + request.runtime.config = {"configurable": {"thread_id": "t1"}} + + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware._invoke_hook", + lambda *_args, **_kwargs: SubagentStartDecision( + event=HookEvent.SUBAGENT_START, + continue_processing=False, + stop_reason="no subagents", + ), + ) + + handler = MagicMock() + blocked = middleware.wrap_tool_call(request, handler) + assert isinstance(blocked, ToolMessage) + assert blocked.status == "error" + assert "no subagents" in str(blocked.content) + handler.assert_not_called() + + async def test_fulfill_hook_invocation_runs_engine(tmp_path: Path) -> None: config_dir = tmp_path / "config" config_dir.mkdir() diff --git a/libs/code/tests/unit_tests/test_non_interactive.py b/libs/code/tests/unit_tests/test_non_interactive.py index 8f24166c55..3a51666334 100644 --- a/libs/code/tests/unit_tests/test_non_interactive.py +++ b/libs/code/tests/unit_tests/test_non_interactive.py @@ -5,6 +5,7 @@ import signal import sys from collections.abc import AsyncIterator, Iterator, Sequence +from pathlib import Path from types import SimpleNamespace from typing import TYPE_CHECKING, Any from unittest.mock import AsyncMock, MagicMock, call, patch @@ -1310,6 +1311,74 @@ async def test_run_agent_loop_passes_thread_id_context(self) -> None: _, kwargs = agent.astream.call_args assert kwargs["context"]["thread_id"] == "t1" + async def test_run_agent_loop_defaults_project_hooks_untrusted( + self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch + ) -> None: + """Headless mode does not load project hooks without explicit trust.""" + monkeypatch.chdir(tmp_path) + project_hooks = tmp_path / ".deepagents" + project_hooks.mkdir() + (project_hooks / "hooks.json").write_text( + '{"hooks":{"Stop":[{"hooks":[{"type":"command","command":"echo x"}]}]}}', + encoding="utf-8", + ) + agent = MagicMock() + agent.astream = MagicMock(return_value=_async_iter([])) + console = Console(quiet=True) + file_op_tracker = MagicMock() + config: RunnableConfig = {"configurable": {"thread_id": "t1"}} + + with patch( + "deepagents_code.client.non_interactive.dispatch_hook", + new_callable=AsyncMock, + ): + await _run_agent_loop( + agent, + "task", + config, + console, + file_op_tracker, + quiet=True, + ) + + _, kwargs = agent.astream.call_args + # Untrusted workspaces omit project Stop handlers from the gate. + assert "Stop" not in (kwargs["context"].get("hooks_server_events") or []) + + async def test_run_agent_loop_trusts_project_hooks_when_opted_in( + self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch + ) -> None: + """`--trust-project-hooks` loads repository hook handlers.""" + monkeypatch.chdir(tmp_path) + project_hooks = tmp_path / ".deepagents" + project_hooks.mkdir() + (project_hooks / "hooks.json").write_text( + '{"hooks":{"Stop":[{"hooks":[{"type":"command","command":"echo x"}]}]}}', + encoding="utf-8", + ) + agent = MagicMock() + agent.astream = MagicMock(return_value=_async_iter([])) + console = Console(quiet=True) + file_op_tracker = MagicMock() + config: RunnableConfig = {"configurable": {"thread_id": "t1"}} + + with patch( + "deepagents_code.client.non_interactive.dispatch_hook", + new_callable=AsyncMock, + ): + await _run_agent_loop( + agent, + "task", + config, + console, + file_op_tracker, + quiet=True, + trust_project_hooks=True, + ) + + _, kwargs = agent.astream.call_args + assert "Stop" in (kwargs["context"].get("hooks_server_events") or []) + async def test_raises_after_user_limit(self) -> None: """HITLIterationLimitError is raised after max_turns HITL iterations.""" agent = _make_looping_agent() From e7eb3441b4d7a04628b0647d366c19f8bbbf90b8 Mon Sep 17 00:00:00 2001 From: Johannes du Plessis Date: Fri, 24 Jul 2026 08:32:37 -0700 Subject: [PATCH 07/14] test(code): expect ServerHooksMiddleware on subagent stacks Co-authored-by: Cursor --- libs/code/tests/unit_tests/test_agent.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/libs/code/tests/unit_tests/test_agent.py b/libs/code/tests/unit_tests/test_agent.py index 018fa5cbd4..149953a2cb 100644 --- a/libs/code/tests/unit_tests/test_agent.py +++ b/libs/code/tests/unit_tests/test_agent.py @@ -3522,6 +3522,7 @@ def test_subagent_middleware_combines_shell_and_configurable_model( """ from deepagents_code.agent import ShellAllowListMiddleware from deepagents_code.configurable_model import ConfigurableModelMiddleware + from deepagents_code.hooks.server_middleware import ServerHooksMiddleware mock_settings = self._build_mock_settings(tmp_path) mock_agent = Mock() @@ -3581,7 +3582,9 @@ def test_subagent_middleware_combines_shell_and_configurable_model( assert middleware_types == [ ConfigurableModelMiddleware, ShellAllowListMiddleware, + ServerHooksMiddleware, ], f"Unexpected middleware on subagent {name!r}: {middleware_types}" + assert subagents_by_name[name]["middleware"][-1]._emit_stop is False pinned = subagents_by_name["pinned"] assert pinned["model"] == "anthropic:claude-haiku-4-5" @@ -3592,6 +3595,13 @@ def test_subagent_middleware_combines_shell_and_configurable_model( assert not any( isinstance(mw, ConfigurableModelMiddleware) for mw in pinned_middleware ), "Pinned subagent must not gain configurable model middleware" + assert any(isinstance(mw, ServerHooksMiddleware) for mw in pinned_middleware), ( + "Pinned subagent should wrap tools with server hooks" + ) + hooks_mw = next( + mw for mw in pinned_middleware if isinstance(mw, ServerHooksMiddleware) + ) + assert hooks_mw._emit_stop is False def test_subagents_get_managed_memory_guard_when_memory_enabled( self, tmp_path: Path From 8a15a162b62954a3fef9a4fe4c90947203660165 Mon Sep 17 00:00:00 2001 From: Johannes du Plessis Date: Fri, 24 Jul 2026 08:30:08 -0700 Subject: [PATCH 08/14] feat(code): integrate Hooks v2 client lifecycle events Wire SessionStart/End, Notification, PermissionRequest, and compact SessionStart through the client HooksRuntime, and keep project hooks behind `--trust-project-hooks` after rebasing onto the server-lifecycle hardening. Co-authored-by: Cursor --- .gitignore | 3 + libs/code/deepagents_code/app.py | 209 +++++++++- .../deepagents_code/client/non_interactive.py | 338 +++++++++++++++-- .../deepagents_code/hooks/client_lifecycle.py | 317 ++++++++++++++++ .../deepagents_code/hooks/models/domain.py | 8 + libs/code/deepagents_code/hooks/projection.py | 13 +- libs/code/deepagents_code/hooks/runtime.py | 9 + .../deepagents_code/tui/textual_adapter.py | 356 +++++++++++++++++- .../unit_tests/hooks/test_client_lifecycle.py | 320 ++++++++++++++++ .../tests/unit_tests/hooks/test_engine.py | 5 +- libs/code/tests/unit_tests/test_app.py | 43 +++ .../tests/unit_tests/test_non_interactive.py | 57 +++ .../unit_tests/tui/test_textual_adapter.py | 64 ++++ 13 files changed, 1699 insertions(+), 43 deletions(-) create mode 100644 libs/code/deepagents_code/hooks/client_lifecycle.py create mode 100644 libs/code/tests/unit_tests/hooks/test_client_lifecycle.py diff --git a/.gitignore b/.gitignore index 27be2ced1e..6f8e48950f 100644 --- a/.gitignore +++ b/.gitignore @@ -228,6 +228,9 @@ __marimo__/ # LangGraph .langgraph_api +# Deep Agents local runtime state +.deepagents/ + #claude .claude diff --git a/libs/code/deepagents_code/app.py b/libs/code/deepagents_code/app.py index b470c79225..706587b27d 100644 --- a/libs/code/deepagents_code/app.py +++ b/libs/code/deepagents_code/app.py @@ -594,6 +594,12 @@ class _ConfigWriteResult: from deepagents_code.config_manifest import CursorStyle from deepagents_code.event_bus import EventSource, ExternalEvent from deepagents_code.goal_rubric import GoalCreateRequest, GoalCriteriaRequest + from deepagents_code.hooks.client_lifecycle import ( + ClientHookContext, + ClientHookService, + ) + from deepagents_code.hooks.models.domain import SessionEndCause, SessionStartCause + from deepagents_code.hooks.runtime import HooksRuntime from deepagents_code.mcp_tools import MCPServerInfo from deepagents_code.model_config import MissingProviderPackageError from deepagents_code.plugins.models import ( @@ -2261,8 +2267,8 @@ def __init__( # Assign the backing field directly: the setter reads `self._thread_id` # to detect a thread change, and it isn't set yet. self._thread_id = thread_id or _new_thread_id() - # Optional session-scoped Hooks v2 client runtime. - self.hooks_runtime = None + self.hooks_runtime: HooksRuntime | None = None + self.client_hooks: ClientHookService | None = None @property def auto_approve(self) -> bool: @@ -2885,6 +2891,7 @@ def __init__( Resolved into a concrete `_lc_thread_id` by `_resolve_resume_thread` during background startup. """ + self._initial_resume_requested = resume_thread is not None self._resume_thread_resolved_event = asyncio.Event() """Set once `-r` resume resolution has completed or is unnecessary.""" @@ -3548,6 +3555,8 @@ def __init__( Lazily constructed by the session-init worker so we don't block startup on it. """ + self._session_state_ready = asyncio.Event() + self._session_init_started = False self._startup_task: asyncio.Task[None] | None = None """Startup task reference (set in on_mount).""" @@ -4200,6 +4209,7 @@ async def _post_paint_init(self) -> None: group="startup-skill-discovery", ) + self._session_init_started = True self.run_worker(self._init_session_state, exclusive=True, group="session-init") # Server startup (model creation + server process) @@ -4279,6 +4289,7 @@ async def _post_paint_init(self) -> None: async def _init_session_state(self) -> None: """Create session state in a thread (imports deepagents_code.sessions).""" + self._session_init_started = True def _create() -> TextualSessionState: from pathlib import Path @@ -4310,14 +4321,115 @@ def _create() -> TextualSessionState: severity="error", timeout=10, ) + self._session_state_ready.set() return # A user can change the approval mode while session construction runs # in the worker thread. Re-read the app-owned selection on the event # loop so the newly assigned state cannot overwrite that newer choice. session_state.approval_mode = self._approval_mode + if session_state.hooks_runtime is not None: + from deepagents_code.hooks.client_lifecycle import ClientHookService + + session_state.client_hooks = ClientHookService( + session_state.hooks_runtime, + notice=lambda message: self.notify(message, markup=False), + ) self._session_state = session_state + self._session_state_ready.set() await self._auto_accept_pending_goal_rubric() + async def _refresh_client_hooks_runtime(self) -> None: + from pathlib import Path + + from deepagents_code.hooks.client_lifecycle import ClientHookService + from deepagents_code.hooks.runtime import HooksRuntime + + state = self._session_state + if state is None: + return + try: + runtime = await asyncio.to_thread( + HooksRuntime.create, + cwd=Path(self._cwd), + workspace_trusted=False, + ) + except Exception: + logger.exception("Failed to refresh HooksRuntime; hooks disabled") + state.hooks_runtime = None + state.client_hooks = None + return + state.hooks_runtime = runtime + state.client_hooks = ClientHookService( + runtime, + notice=lambda message: self.notify(message, markup=False), + ) + + def _client_hook_context( + self, *, thread_id: str | None = None + ) -> ClientHookContext | None: + from deepagents_code.hooks.client_lifecycle import ClientHookContext + + state = self._session_state + if state is None: + return None + return ClientHookContext.create( + thread_id=thread_id or state.thread_id, + approval_mode=state.approval_mode, + prompt_id=state.turn_id, + ) + + def _client_hook_service(self) -> ClientHookService | None: + from deepagents_code.hooks.client_lifecycle import ClientHookService + + state = self._session_state + if state is None or not isinstance(state.client_hooks, ClientHookService): + return None + return state.client_hooks + + async def _run_session_start_hook(self, cause: SessionStartCause) -> bool: + from deepagents_code.config import settings + from deepagents_code.hooks.models.domain import HookEvent + + service = self._client_hook_service() + if service is None or not service.has_handlers(HookEvent.SESSION_START): + return True + context = self._client_hook_context() + if context is None: + return True + try: + decision = await service.session_start( + context, + cause, + model=settings.model_name or None, + ) + except Exception: + logger.warning("SessionStart hook invocation failed", exc_info=True) + return True + if decision.continue_processing: + return True + message = decision.stop_reason or "Session start was stopped by a hook." + await self._mount_message(AppMessage(message)) + return False + + async def _run_session_end_hook( + self, + cause: SessionEndCause, + *, + thread_id: str | None = None, + ) -> None: + from deepagents_code.hooks.models.domain import HookEvent + + service = self._client_hook_service() + if service is None or not service.has_handlers(HookEvent.SESSION_END): + return + context = self._client_hook_context(thread_id=thread_id) + if context is None: + return + try: + await service.session_end(context, cause) + except Exception: + logger.warning("SessionEnd hook invocation failed", exc_info=True) + async def _ensure_managed_ripgrep(self) -> bool: """Install the managed `rg` and prepend it to `PATH`, exactly once. @@ -4652,6 +4764,7 @@ async def _resolve_resume_thread(self) -> None: candidate = await get_most_recent(agent_filter) if not candidate: self._lc_thread_id = generate_thread_id() + self._initial_resume_requested = False self._resuming = False self._sync_status_connection() if agent_filter: @@ -4665,6 +4778,7 @@ async def _resolve_resume_thread(self) -> None: else: # Thread not found — notify + fall back to new thread self._lc_thread_id = generate_thread_id() + self._initial_resume_requested = False self._resuming = False self._sync_status_connection() similar = await find_similar_threads(resume) @@ -4709,6 +4823,7 @@ async def _resolve_resume_thread(self) -> None: # User declined the resume: start a fresh session and skip the # agent/model adoption below so it inherits the launch default. self._lc_thread_id = generate_thread_id() + self._initial_resume_requested = False self._resuming = False self._sync_status_connection() self.notify( @@ -4733,6 +4848,7 @@ async def _resolve_resume_thread(self) -> None: except Exception: logger.exception("Failed to resolve resume thread %r", resume) self._lc_thread_id = generate_thread_id() + self._initial_resume_requested = False self._resuming = False self._sync_status_connection() self.notify( @@ -8080,10 +8196,22 @@ async def _run_session_start_sequence(self) -> None: self._schedule_session_start_after_launch_init(launch_init_task) return + if self._session_state is None and self._session_init_started: + await self._session_state_ready.wait() + self._initial_session_started = True self._startup_sequence_running = True initial_submitted = False try: + from deepagents_code.hooks.models.domain import SessionStartCause + + start_cause = ( + SessionStartCause.RESUME + if self._initial_resume_requested + else SessionStartCause.STARTUP + ) + if not await self._run_session_start_hook(start_cause): + return should_load_history = bool(self._lc_thread_id and self._agent) and ( self._resume_thread_intent is not None or not self._has_initial_submission() @@ -12440,8 +12568,14 @@ async def _handle_command(self, command: str) -> None: ): await self._handle_rubric_command(command) elif cmd in {"/clear", "/force-clear"}: + from deepagents_code.hooks.models.domain import ( + SessionEndCause, + SessionStartCause, + ) + if cmd == "/force-clear": self._force_interrupt_active_work() + await self._run_session_end_hook(SessionEndCause.CLEAR) self._pending_messages.clear() self._queued_widgets.clear() self._sync_status_queued() @@ -12520,6 +12654,9 @@ async def _handle_command(self, command: str) -> None: thread_id=previous_thread_id, suffix=resume_hint, ) + await self._refresh_client_hooks_runtime() + if not await self._run_session_start_hook(SessionStartCause.CLEAR): + return elif cmd == "/copy": await self._mount_message(UserMessage(command)) # Reverse-scan for the newest assistant message that has finished @@ -13481,6 +13618,10 @@ async def _handle_offload(self) -> None: f"Context: {before} → {after} tokens " f"({pct}% decrease), {messages_kept} messages kept." ) + from deepagents_code.hooks.models.domain import SessionStartCause + + if not await self._run_session_start_hook(SessionStartCause.COMPACT): + return if archive_path: from deepagents_code.offload import offload_storage_is_ephemeral @@ -16589,6 +16730,15 @@ def exit( # `finally` in `run_textual_app`; `ServerProcess.stop()` is idempotent # and serialized, so the two callers never race or double-clean. server_proc = self._server_proc + from deepagents_code.hooks.models.domain import HookEvent + + client_hook_service = self._client_hook_service() + has_client_session_hooks = ( + client_hook_service is not None + and client_hook_service.has_handlers(HookEvent.SESSION_END) + ) + if has_client_session_hooks: + session_end_payload = None if ( should_wait_for_agent @@ -16596,6 +16746,7 @@ def exit( or should_drain_hooks or server_proc is not None or session_end_payload is not None + or has_client_session_hooks ): refreshed: asyncio.Event | None = None if should_wait_for_agent or should_drain_hooks: @@ -16655,6 +16806,15 @@ async def _dispatch_session_end() -> None: ) session_end_task = asyncio.ensure_future(_dispatch_session_end()) + client_session_end_task: asyncio.Task[None] | None = None + if has_client_session_hooks: + from deepagents_code.hooks.models.domain import SessionEndCause + + client_session_end_task = asyncio.ensure_future( + self._run_session_end_hook( + SessionEndCause.PROMPT_INPUT_EXIT, + ) + ) async def _drain_hooks() -> None: phase_start = time.monotonic() @@ -16839,6 +16999,14 @@ async def _finish_restart() -> None: "(force-quit); dispatch may not have completed", exc_info=True, ) + if client_session_end_task is not None: + try: + await client_session_end_task + except BaseException: + logger.debug( + "SessionEnd await interrupted during teardown", + exc_info=True, + ) logger.debug( "Teardown total took %.3fs", time.monotonic() - teardown_start, @@ -18346,6 +18514,12 @@ def _build_agent(url: str) -> Any: # noqa: ANN401 # see docstring self._update_status("") if self._session_state: + from deepagents_code.hooks.models.domain import SessionEndCause + + await self._run_session_end_hook( + SessionEndCause.OTHER, + thread_id=previous_thread_id, + ) new_thread_id = self._session_state.reset_thread() self._lc_thread_id = new_thread_id self._update_welcome_banner( @@ -18454,6 +18628,11 @@ def _build_agent(url: str) -> Any: # noqa: ANN401 # see docstring exc_info=True, ) self._sync_status_connection() + from deepagents_code.hooks.models.domain import SessionStartCause + + await self._refresh_client_hooks_runtime() + if not await self._run_session_start_hook(SessionStartCause.CLEAR): + return # Refresh skills so /skill: autocomplete reflects the new agent's # SKILL.md files. @@ -22151,6 +22330,15 @@ async def _resume_thread(self, thread_id: str) -> None: if await asyncio.to_thread(self._cwd_paths_equal, self._cwd, prev_cwd): await self._mount_message(AppMessage(f"Already on thread: {thread_id}")) else: + from deepagents_code.hooks.models.domain import ( + SessionEndCause, + SessionStartCause, + ) + + await self._run_session_end_hook(SessionEndCause.RESUME) + await self._refresh_client_hooks_runtime() + if not await self._run_session_start_hook(SessionStartCause.RESUME): + return await self._mount_message( AppMessage(f"Switched to thread directory: {self._cwd}"), ) @@ -22178,10 +22366,21 @@ async def _resume_thread(self, thread_id: str) -> None: self._chat_input.set_cursor_active(active=False) prefetched_payload: _ThreadHistoryPayload | None = None + outgoing_ended = False try: self._update_status(f"Loading thread: {thread_id}") await self._set_spinner("Loading thread") prefetched_payload = await self._fetch_thread_history_data(thread_id) + from deepagents_code.hooks.models.domain import ( + SessionEndCause, + SessionStartCause, + ) + + await self._run_session_end_hook( + SessionEndCause.RESUME, + thread_id=prev_session_thread, + ) + outgoing_ended = True # Clear conversation (similar to /clear, without creating a new thread) await self._set_spinner(None) @@ -22222,6 +22421,9 @@ async def _resume_thread(self, thread_id: str) -> None: # thread". Set only after the last statement that can raise, so a # failed switch (handled below) never leaves a stale pointer. self._session_state.previous_thread_id = prev_session_thread + await self._refresh_client_hooks_runtime() + if not await self._run_session_start_hook(SessionStartCause.RESUME): + return except Exception as exc: if prefetched_payload is None: logger.exception("Failed to prefetch history for thread %s", thread_id) @@ -22258,6 +22460,9 @@ async def _resume_thread(self, thread_id: str) -> None: "switch to %s" ) logger.warning(msg, thread_id, exc_info=True) + if outgoing_ended: + await self._refresh_client_hooks_runtime() + await self._run_session_start_hook(SessionStartCause.RESUME) error_message = f"Failed to switch to thread {thread_id}: {exc}." if rollback_restore_failed: error_message += " Previous thread history could not be restored." diff --git a/libs/code/deepagents_code/client/non_interactive.py b/libs/code/deepagents_code/client/non_interactive.py index c538d660fe..c2b1f23558 100644 --- a/libs/code/deepagents_code/client/non_interactive.py +++ b/libs/code/deepagents_code/client/non_interactive.py @@ -26,7 +26,7 @@ import threading import time from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any, cast +from typing import TYPE_CHECKING, Any, NoReturn, cast from langchain.agents.middleware.human_in_the_loop import ActionRequest, HITLRequest from langchain_core.messages import AIMessage, ToolMessage @@ -83,10 +83,13 @@ if TYPE_CHECKING: from asyncio.subprocess import Process from pathlib import Path + from uuid import UUID from deepagents import FsToolName from langchain_core.runnables import RunnableConfig + from deepagents_code.approval_mode import ApprovalMode + from deepagents_code.hooks.client_lifecycle import ClientHookService from deepagents_code.hooks.runtime import HooksRuntime logger = logging.getLogger(__name__) @@ -96,6 +99,15 @@ class HITLIterationLimitError(RuntimeError): """Raised when the HITL interrupt loop exceeds `_MAX_HITL_ITERATIONS` rounds.""" +def _raise_hitl_iteration_limit(message: str) -> NoReturn: + """Raise the bounded-turn failure outside the stream-control try block. + + Raises: + HITLIterationLimitError: Always, with the supplied message. + """ + raise HITLIterationLimitError(message) + + _HITL_REQUEST_ADAPTER = TypeAdapter(HITLRequest) _STREAM_CHUNK_LENGTH = 3 @@ -365,6 +377,12 @@ class StreamState: hooks_runtime: HooksRuntime | None = None """Optional session-scoped HooksRuntime used to fulfill server hook interrupts.""" + client_hooks: ClientHookService | None = None + """Client-owned lifecycle facade for this headless session.""" + + summarization_observed: bool = False + """Whether the current stream crossed a compaction boundary.""" + interrupt_occurred: bool = False """Flag indicating whether any HITL interrupt was received during the current stream pass.""" @@ -431,6 +449,7 @@ def _process_interrupts( console: Rich console for user-visible warnings. """ from deepagents_code.hooks.interrupt import is_hook_interrupt_payload + from deepagents_code.hooks.models.domain import HookEvent interrupts = data["__interrupt__"] if interrupts: @@ -461,7 +480,10 @@ def _process_interrupts( continue state.pending_interrupts[interrupt_obj.id] = validated_request state.interrupt_occurred = True - dispatch_hook_fire_and_forget("input.required", {}) + if state.client_hooks is None or not state.client_hooks.has_handlers( + HookEvent.NOTIFICATION + ): + dispatch_hook_fire_and_forget("input.required", {}) def _process_ai_message( @@ -618,6 +640,7 @@ def _process_message_chunk( # conversation history for the LLM. These are internal bookkeeping and # should not be rendered to the user. if metadata and metadata.get("lc_source") == "summarization": + state.summarization_observed = True return if isinstance(message_obj, AIMessage): @@ -971,7 +994,11 @@ async def _fulfill_pending_hook_interrupts(state: StreamState) -> None: ) -def _process_hitl_interrupts(state: StreamState, console: Console) -> None: +async def _process_hitl_interrupts( + state: StreamState, + console: Console, + thread_id: str, +) -> None: """Iterate over pending HITL interrupts and build approval/rejection responses. After processing, `state.pending_interrupts` is cleared and decisions @@ -980,16 +1007,98 @@ def _process_hitl_interrupts(state: StreamState, console: Console) -> None: Args: state: Stream state containing the pending interrupts to process. console: Rich console for status output. + thread_id: Active conversation thread. + + Raises: + ClientHookStopError: If a hook stops or interrupts approval. """ current_interrupts = dict(state.pending_interrupts) state.pending_interrupts.clear() + from deepagents_code.approval_mode import ApprovalMode + from deepagents_code.hooks.client_lifecycle import ( + ClientHookContext, + ClientHookStopError, + ) + from deepagents_code.hooks.models.domain import ( + DcodeNotificationKind, + ToolCallData, + ) + + context = ClientHookContext.create( + thread_id=thread_id, + approval_mode=ApprovalMode.MANUAL, + ) for interrupt_id, hitl_request in current_interrupts.items(): - decisions = [ - _make_hitl_decision(action_request, console) - for action_request in hitl_request["action_requests"] - ] - state.hitl_response[interrupt_id] = {"decisions": decisions} + action_requests = hitl_request["action_requests"] + decisions: list[dict[str, str] | None] = [] + for index, action_request in enumerate(action_requests): + try: + hook_decision = ( + await state.client_hooks.permission_request( + context, + ToolCallData( + id=f"{interrupt_id}:{index}", + name=action_request.get("name", ""), + args=action_request.get("args", {}), + ), + ) + if state.client_hooks is not None + else None + ) + except Exception: + logger.warning( + "PermissionRequest hook invocation failed", + exc_info=True, + ) + hook_decision = None + if hook_decision is None: + decisions.append(None) + continue + if not hook_decision.continue_processing: + reason = hook_decision.stop_reason or "Permission stopped by hook" + raise ClientHookStopError(reason) + permission = hook_decision.permission + if permission.behavior == "allow": + decisions.append({"type": "approve"}) + elif permission.behavior == "deny": + denied = {"type": "reject"} + if permission.reason: + denied["message"] = permission.reason + if permission.interrupt: + raise ClientHookStopError( + permission.reason or "Permission interrupted by hook" + ) + decisions.append(denied) + else: + decisions.append(None) + + if ( + any(decision is None for decision in decisions) + and state.client_hooks is not None + ): + try: + await state.client_hooks.notification( + context, + DcodeNotificationKind.PERMISSION_REQUIRED, + "Permission required", + ) + except ClientHookStopError: + raise + except Exception: + logger.warning("Notification hook invocation failed", exc_info=True) + resolved: list[dict[str, str]] = [] + for decision, action_request in zip( + decisions, + action_requests, + strict=True, + ): + resolved.append( + decision + if decision is not None + else _make_hitl_decision(action_request, console) + ) + state.hitl_response[interrupt_id] = {"decisions": resolved} async def _stream_agent( @@ -1084,6 +1193,9 @@ async def _run_agent_loop( rubric: str | None = None, show_rubric_iterations: bool = False, trust_project_hooks: bool = False, + hooks_runtime: HooksRuntime | None = None, + approval_mode: ApprovalMode | None = None, + prompt_id: UUID | None = None, ) -> None: """Run the agent and handle HITL interrupts until the task completes. @@ -1121,9 +1233,12 @@ async def _run_agent_loop( Defaults to `False` so untrusted checkouts cannot execute repository hooks in CI without an explicit opt-in. + hooks_runtime: Preloaded session runtime, when available. + approval_mode: Effective client approval policy. Defaults to manual. + prompt_id: Stable identifier for the headless turn. Raises: - HITLIterationLimitError: If the effective turn limit is exceeded. + ClientHookStopError: If a client-owned hook stops processing. """ spinner = None if quiet else _ConsoleSpinner(console) state = StreamState( @@ -1150,20 +1265,87 @@ async def _run_agent_loop( from pathlib import Path + from deepagents_code.approval_mode import ApprovalMode + from deepagents_code.hooks.client_lifecycle import ( + ClientHookContext, + ClientHookService, + ClientHookStopError, + ) from deepagents_code.hooks.context import apply_hooks_context + from deepagents_code.hooks.models.domain import ( + DcodeNotificationKind, + HookEvent, + SessionEndCause, + SessionStartCause, + ) from deepagents_code.hooks.runtime import HooksRuntime - try: - # Project hooks require an explicit opt-in, matching `--trust-project-mcp`. - hooks_runtime = HooksRuntime.create( - cwd=Path.cwd(), - workspace_trusted=trust_project_hooks, - ) - except Exception: - logger.exception("Failed to create HooksRuntime; server hooks disabled") - hooks_runtime = None - apply_hooks_context(context, hooks_runtime) + resolved_approval_mode = approval_mode or ApprovalMode.MANUAL + if hooks_runtime is None: + try: + # Project hooks require an explicit opt-in, matching `--trust-project-mcp`. + hooks_runtime = HooksRuntime.create( + cwd=Path.cwd(), + workspace_trusted=trust_project_hooks, + ) + except Exception: + logger.exception("Failed to create HooksRuntime; server hooks disabled") + hooks_runtime = None + apply_hooks_context( + context, + hooks_runtime, + prompt_id=str(prompt_id) if prompt_id is not None else None, + ) + context["approval_mode"] = resolved_approval_mode.value + context["auto_approve"] = resolved_approval_mode is ApprovalMode.YOLO state.hooks_runtime = hooks_runtime + state.client_hooks = ( + ClientHookService( + hooks_runtime, + notice=lambda notice: console.print(Text(notice), highlight=False), + ) + if hooks_runtime is not None + else None + ) + + client_context = ClientHookContext.create( + thread_id=thread_id, + approval_mode=resolved_approval_mode, + prompt_id=prompt_id, + ) + if state.client_hooks is not None: + try: + start_decision = await state.client_hooks.session_start( + client_context, + SessionStartCause.STARTUP, + model=settings.model_name or None, + ) + except Exception: + logger.warning("SessionStart hook invocation failed", exc_info=True) + start_decision = None + if start_decision is None: + session_context = () + else: + session_context = state.client_hooks.take_session_context(thread_id) + if start_decision is not None and not start_decision.continue_processing: + reason = start_decision.stop_reason or "Session start stopped by hook" + try: + await state.client_hooks.session_end( + client_context, + SessionEndCause.OTHER, + ) + except Exception: + logger.warning("SessionEnd hook invocation failed", exc_info=True) + raise ClientHookStopError(reason) + if session_context: + messages = stream_input["messages"] + messages.insert( + 0, + { + "role": "system", + "content": "\n\n".join(session_context), + }, + ) await dispatch_hook("session.start", {"thread_id": thread_id}) @@ -1174,6 +1356,26 @@ async def _run_agent_loop( await _stream_agent( agent, stream_input, config, state, console, file_op_tracker, context ) + if state.summarization_observed and state.client_hooks is not None: + try: + compact_decision = await state.client_hooks.session_start( + client_context, + SessionStartCause.COMPACT, + model=settings.model_name or None, + ) + except Exception: + logger.warning( + "Compact SessionStart hook invocation failed", + exc_info=True, + ) + else: + if not compact_decision.continue_processing: + reason = ( + compact_decision.stop_reason + or "Compact session start stopped by hook" + ) + raise ClientHookStopError(reason) + state.summarization_observed = False # The internal default applies when --max-turns is omitted, guarding # against unbounded runaway loops in scripts that forgot to set one. @@ -1195,18 +1397,48 @@ async def _run_agent_loop( "The agent may be stuck retrying rejected commands. " "Increase --max-turns or break the task into smaller steps." ) - raise HITLIterationLimitError(msg) + _raise_hitl_iteration_limit(msg) turns += 1 state.interrupt_occurred = False state.hitl_response.clear() state.hook_response.clear() await _fulfill_pending_hook_interrupts(state) - _process_hitl_interrupts(state, console) + await _process_hitl_interrupts(state, console, thread_id) resume_payload = {**state.hook_response, **state.hitl_response} stream_input = Command(resume=resume_payload) await _stream_agent( agent, stream_input, config, state, console, file_op_tracker, context ) + if state.summarization_observed and state.client_hooks is not None: + try: + compact_decision = await state.client_hooks.session_start( + client_context, + SessionStartCause.COMPACT, + model=settings.model_name or None, + ) + except Exception: + logger.warning( + "Compact SessionStart hook invocation failed", + exc_info=True, + ) + else: + if not compact_decision.continue_processing: + reason = ( + compact_decision.stop_reason + or "Compact session start stopped by hook" + ) + raise ClientHookStopError(reason) + state.summarization_observed = False + except BaseException: + if state.client_hooks is not None: + try: + await state.client_hooks.session_end( + client_context, + SessionEndCause.OTHER, + ) + except Exception: + logger.warning("SessionEnd hook invocation failed", exc_info=True) + raise finally: # Close out any `tool.use` with no matching `ToolMessage` — e.g. a stream # aborted by a provider error mid-tool. On a clean run every id was @@ -1275,8 +1507,34 @@ async def _run_agent_loop( console.print("[green]✓ Task completed[/green]") print_usage_table(state.stats, wall_time, console) - await dispatch_hook("task.complete", {"thread_id": thread_id}) - await dispatch_hook("session.end", {"thread_id": thread_id}) + if state.client_hooks is not None: + notification_stop: ClientHookStopError | None = None + try: + await state.client_hooks.notification( + client_context, + DcodeNotificationKind.AGENT_COMPLETED, + "Agent completed", + ) + except ClientHookStopError as exc: + notification_stop = exc + except Exception: + logger.warning("Notification hook invocation failed", exc_info=True) + if not state.client_hooks.has_handlers(HookEvent.NOTIFICATION): + await dispatch_hook("task.complete", {"thread_id": thread_id}) + try: + await state.client_hooks.session_end( + client_context, + SessionEndCause.PROMPT_INPUT_EXIT, + ) + except Exception: + logger.warning("SessionEnd hook invocation failed", exc_info=True) + if not state.client_hooks.has_handlers(HookEvent.SESSION_END): + await dispatch_hook("session.end", {"thread_id": thread_id}) + if notification_stop is not None: + raise notification_stop + else: + await dispatch_hook("task.complete", {"thread_id": thread_id}) + await dispatch_hook("session.end", {"thread_id": thread_id}) def _build_non_interactive_header( @@ -1668,6 +1926,24 @@ def discover_all_skills() -> tuple[list[ExtendedSkillMetadata], list[Path]]: logger.warning("MCP metadata preload task creation failed", exc_info=True) try: + from pathlib import Path + + from deepagents_code.approval_mode import ApprovalMode + from deepagents_code.hooks.models.domain import HookEvent + from deepagents_code.hooks.runtime import HooksRuntime + + try: + hooks_runtime = HooksRuntime.create( + cwd=Path.cwd(), + workspace_trusted=trust_project_hooks, + ) + except Exception: + logger.exception("Failed to create HooksRuntime; hooks disabled") + hooks_runtime = None + has_permission_hooks = bool( + hooks_runtime is not None + and HookEvent.PERMISSION_REQUEST in hooks_runtime.configured_events() + ) enable_shell = bool(settings.shell_allow_list) shell_is_unrestricted = isinstance( settings.shell_allow_list, type(SHELL_ALLOW_ALL) @@ -1675,8 +1951,14 @@ def discover_all_skills() -> tuple[list[ExtendedSkillMetadata], list[Path]]: # Currently, non-shell tools have no HITL handler in non-interactive # mode, so interrupting on them just fragments LangSmith traces # without adding value. Gate only shell execution via middleware. - use_auto_approve = not enable_shell or shell_is_unrestricted - use_interrupt_shell_only = enable_shell and not shell_is_unrestricted + requested_auto_approve = not enable_shell or shell_is_unrestricted + use_auto_approve = requested_auto_approve and not has_permission_hooks + use_interrupt_shell_only = ( + enable_shell and not shell_is_unrestricted and not has_permission_hooks + ) + approval_mode = ( + ApprovalMode.YOLO if requested_auto_approve else ApprovalMode.MANUAL + ) # Extract the concrete allow-list to forward to the server subprocess. # settings.shell_allow_list is already validated at this point. restrictive_allow_list: list[str] | None = ( @@ -1693,11 +1975,12 @@ def discover_all_skills() -> tuple[list[ExtendedSkillMetadata], list[Path]]: from deepagents_code.config import build_stream_config + turn_id = uuid4() config: RunnableConfig = build_stream_config( thread_id, assistant_id, sandbox_type=sandbox_type, - turn_id=str(uuid4()), + turn_id=str(turn_id), turn_number=1, auto_approve=use_auto_approve, ) @@ -1764,6 +2047,9 @@ def discover_all_skills() -> tuple[list[ExtendedSkillMetadata], list[Path]]: rubric=rubric, show_rubric_iterations=rubric_max_iterations is not None, trust_project_hooks=trust_project_hooks, + hooks_runtime=hooks_runtime, + approval_mode=approval_mode, + prompt_id=turn_id, ) except KeyboardInterrupt: diff --git a/libs/code/deepagents_code/hooks/client_lifecycle.py b/libs/code/deepagents_code/hooks/client_lifecycle.py new file mode 100644 index 0000000000..aa3134f3f4 --- /dev/null +++ b/libs/code/deepagents_code/hooks/client_lifecycle.py @@ -0,0 +1,317 @@ +"""Client-owned Hooks v2 lifecycle facade.""" + +from __future__ import annotations + +import logging +import sys +from dataclasses import dataclass, field +from typing import TYPE_CHECKING +from uuid import UUID + +from deepagents_code.approval_mode import ApprovalMode +from deepagents_code.hooks.models.domain import ( + DcodeNotification, + DcodeNotificationKind, + HookContext, + HookDecision, + HookDiagnostic, + HookDomainEvent, + HookEvent, + HookInvocation, + NotificationDecision, + NotificationEvent, + PermissionEffect, + PermissionRequestDecision, + PermissionRequestEvent, + SessionEndCause, + SessionEndDecision, + SessionEndEvent, + SessionStartCause, + SessionStartDecision, + SessionStartEvent, + ToolCallData, +) + +if TYPE_CHECKING: + from collections.abc import Callable + from pathlib import Path + from typing import Protocol + + class _ClientHooksRuntime(Protocol): + @property + def cwd(self) -> Path: ... + + def configured_events(self) -> frozenset[HookEvent]: ... + + async def invoke(self, invocation: HookInvocation) -> HookDecision: ... + + +logger = logging.getLogger(__name__) + + +class ClientHookStopError(RuntimeError): + """Raised when a client-owned hook stops lifecycle processing.""" + + +@dataclass(frozen=True, slots=True) +class ClientHookContext: + """Client state required to create a domain hook invocation.""" + + thread_id: str + approval_mode: ApprovalMode + prompt_id: UUID | None = None + + @classmethod + def create( + cls, + *, + thread_id: str, + approval_mode: ApprovalMode | str, + prompt_id: str | UUID | None = None, + ) -> ClientHookContext: + """Build validated client hook context. + + Args: + thread_id: Active conversation thread. + approval_mode: Current client approval policy. + prompt_id: Optional current prompt identifier. + + Returns: + Validated context for client-owned hook events. + """ + approval = ( + approval_mode + if isinstance(approval_mode, ApprovalMode) + else ApprovalMode(approval_mode) + ) + parsed_prompt = ( + prompt_id + if isinstance(prompt_id, UUID) + else UUID(prompt_id) + if prompt_id + else None + ) + return cls( + thread_id=thread_id, + approval_mode=approval, + prompt_id=parsed_prompt, + ) + + +@dataclass(slots=True) +class ClientHookService: + """Execute client-owned events and apply their common side effects.""" + + runtime: _ClientHooksRuntime + notice: Callable[[str], None] | None = None + _session_context: dict[str, list[str]] = field(default_factory=dict) + + async def session_start( + self, + context: ClientHookContext, + cause: SessionStartCause, + *, + model: str | None = None, + ) -> SessionStartDecision: + """Invoke `SessionStart` and retain context for the next model turn. + + Args: + context: Current client session context. + cause: Lifecycle boundary that started the session. + model: Active model identifier when available. + + Returns: + Aggregated session-start decision. + + Raises: + TypeError: If the runtime returns a mismatched decision type. + """ + if not self.has_handlers(HookEvent.SESSION_START): + return SessionStartDecision(event=HookEvent.SESSION_START) + decision = await self._invoke( + context, + SessionStartEvent( + event=HookEvent.SESSION_START, + cause=cause, + model=model, + ), + ) + if not isinstance(decision, SessionStartDecision): + msg = f"Expected SessionStartDecision, got {type(decision).__name__}" + raise TypeError(msg) + if decision.context: + self._session_context.setdefault(context.thread_id, []).extend( + decision.context + ) + return decision + + async def session_end( + self, + context: ClientHookContext, + cause: SessionEndCause, + ) -> SessionEndDecision: + """Invoke `SessionEnd` for the outgoing thread. + + Args: + context: Outgoing client session context. + cause: Reason the session ended. + + Returns: + Aggregated session-end decision. + + Raises: + TypeError: If the runtime returns a mismatched decision type. + """ + if not self.has_handlers(HookEvent.SESSION_END): + self._session_context.pop(context.thread_id, None) + return SessionEndDecision(event=HookEvent.SESSION_END) + decision = await self._invoke( + context, + SessionEndEvent(event=HookEvent.SESSION_END, cause=cause), + ) + if not isinstance(decision, SessionEndDecision): + msg = f"Expected SessionEndDecision, got {type(decision).__name__}" + raise TypeError(msg) + self._session_context.pop(context.thread_id, None) + return decision + + async def permission_request( + self, + context: ClientHookContext, + call: ToolCallData, + ) -> PermissionRequestDecision: + """Invoke `PermissionRequest` before client approval resolution. + + Args: + context: Current client session context. + call: Tool action awaiting approval. + + Returns: + Aggregated permission decision. + + Raises: + TypeError: If the runtime returns a mismatched decision type. + """ + if not self.has_handlers(HookEvent.PERMISSION_REQUEST): + return PermissionRequestDecision( + event=HookEvent.PERMISSION_REQUEST, + permission=PermissionEffect(behavior="none"), + ) + decision = await self._invoke( + context, + PermissionRequestEvent(event=HookEvent.PERMISSION_REQUEST, call=call), + ) + if not isinstance(decision, PermissionRequestDecision): + msg = f"Expected PermissionRequestDecision, got {type(decision).__name__}" + raise TypeError(msg) + return decision + + async def notification( + self, + context: ClientHookContext, + kind: DcodeNotificationKind, + message: str, + *, + title: str | None = None, + ) -> NotificationDecision: + """Invoke one explicitly supported dcode notification event. + + Args: + context: Current client session context. + kind: Supported dcode notification kind. + message: User-facing notification text. + title: Optional notification title. + + Returns: + Aggregated notification decision. + + Raises: + ClientHookStopError: If a handler stops lifecycle processing. + TypeError: If the runtime returns a mismatched decision type. + """ + if not self.has_handlers(HookEvent.NOTIFICATION): + return NotificationDecision(event=HookEvent.NOTIFICATION) + decision = await self._invoke( + context, + NotificationEvent( + event=HookEvent.NOTIFICATION, + notification=DcodeNotification( + type=kind, + message=message, + title=title, + ), + ), + ) + if not isinstance(decision, NotificationDecision): + msg = f"Expected NotificationDecision, got {type(decision).__name__}" + raise TypeError(msg) + if not decision.continue_processing: + reason = decision.stop_reason or "Notification stopped by hook" + raise ClientHookStopError(reason) + return decision + + def take_session_context(self, thread_id: str) -> tuple[str, ...]: + """Consume context accumulated for the thread's next model turn. + + Args: + thread_id: Thread whose pending context should be consumed. + + Returns: + Ordered context strings, removed from the service. + """ + return tuple(self._session_context.pop(thread_id, ())) + + def has_handlers(self, event: HookEvent) -> bool: + """Return whether the runtime has handlers for an event. + + Args: + event: Lifecycle event to inspect. + + Returns: + Whether at least one handler was configured. + """ + return event in self.runtime.configured_events() + + async def _invoke( + self, + context: ClientHookContext, + event: HookDomainEvent, + ) -> HookDecision: + invocation = HookInvocation( + context=HookContext( + thread_id=context.thread_id, + cwd=self.runtime.cwd, + prompt_id=context.prompt_id, + approval_mode=context.approval_mode, + ), + event=event, + ) + decision = await self.runtime.invoke(invocation) + self._apply_common_effects(decision) + return decision + + def _apply_common_effects(self, decision: HookDecision) -> None: + for diagnostic in decision.diagnostics: + _log_diagnostic(diagnostic) + for notice in decision.user_notices: + if self.notice is None: + logger.warning("Hook user notice: %s", notice) + continue + try: + self.notice(notice) + except Exception: + logger.warning("Failed to surface hook user notice", exc_info=True) + for sequence in decision.terminal_sequences: + sys.stdout.write(sequence) + if decision.terminal_sequences: + sys.stdout.flush() + + +def _log_diagnostic(diagnostic: HookDiagnostic) -> None: + message = "Hook diagnostic %s: %s" + if diagnostic.severity == "error": + logger.error(message, diagnostic.code, diagnostic.message) + elif diagnostic.severity == "warning": + logger.warning(message, diagnostic.code, diagnostic.message) + else: + logger.debug(message, diagnostic.code, diagnostic.message) diff --git a/libs/code/deepagents_code/hooks/models/domain.py b/libs/code/deepagents_code/hooks/models/domain.py index 3da136a1bf..bf1b96471a 100644 --- a/libs/code/deepagents_code/hooks/models/domain.py +++ b/libs/code/deepagents_code/hooks/models/domain.py @@ -76,6 +76,14 @@ class SessionEndCause(StrEnum): OTHER = "other" +class DcodeNotificationKind(StrEnum): + """dcode lifecycle notifications with compatible wire mappings.""" + + PERMISSION_REQUIRED = "permission_required" + AGENT_NEEDS_INPUT = "agent_needs_input" + AGENT_COMPLETED = "agent_completed" + + class CompactTrigger(StrEnum): """Reason context compaction was requested.""" diff --git a/libs/code/deepagents_code/hooks/projection.py b/libs/code/deepagents_code/hooks/projection.py index 45f27ab87f..c2bd2fe26e 100644 --- a/libs/code/deepagents_code/hooks/projection.py +++ b/libs/code/deepagents_code/hooks/projection.py @@ -10,6 +10,7 @@ from deepagents_code.approval_mode import ApprovalMode from deepagents_code.hooks.models.adapters import HOOK_WIRE_INPUT_ADAPTER from deepagents_code.hooks.models.domain import ( + DcodeNotificationKind, HookEvent, NotificationEvent, PermissionRequestEvent, @@ -361,9 +362,17 @@ def _permission_mode(mode: ApprovalMode) -> WirePermissionMode: def _notification_type(value: str) -> WireNotificationType: + mappings: dict[str, WireNotificationType] = { + DcodeNotificationKind.PERMISSION_REQUIRED: ( + WireNotificationType.PERMISSION_PROMPT + ), + WireNotificationType.PERMISSION_PROMPT: WireNotificationType.PERMISSION_PROMPT, + DcodeNotificationKind.AGENT_NEEDS_INPUT: WireNotificationType.AGENT_NEEDS_INPUT, + DcodeNotificationKind.AGENT_COMPLETED: WireNotificationType.AGENT_COMPLETED, + } try: - return WireNotificationType(value) - except ValueError as exc: + return mappings[value] + except KeyError as exc: msg = f"Unsupported notification type: {value}" raise ValueError(msg) from exc diff --git a/libs/code/deepagents_code/hooks/runtime.py b/libs/code/deepagents_code/hooks/runtime.py index 351c7a003c..e221983fab 100644 --- a/libs/code/deepagents_code/hooks/runtime.py +++ b/libs/code/deepagents_code/hooks/runtime.py @@ -12,6 +12,7 @@ from deepagents_code.hooks.loading import load_hooks_config from deepagents_code.hooks.models.domain import ( HookDecision, + HookEvent, HookInvocation, SubagentStartEvent, SubagentStopEvent, @@ -103,6 +104,14 @@ def configured_server_events(self) -> tuple[str, ...]: sorted(event.value for event in self.snapshot.configured_server_events()) ) + def configured_events(self) -> frozenset[HookEvent]: + """Return every event with at least one configured handler. + + Returns: + Immutable configured event set. + """ + return self.snapshot.configured_events() + def append_messages( self, thread_id: str, diff --git a/libs/code/deepagents_code/tui/textual_adapter.py b/libs/code/deepagents_code/tui/textual_adapter.py index abda43282f..af5b09a6f6 100644 --- a/libs/code/deepagents_code/tui/textual_adapter.py +++ b/libs/code/deepagents_code/tui/textual_adapter.py @@ -13,11 +13,12 @@ import httpx if TYPE_CHECKING: - from collections.abc import Awaitable, Callable, Iterable, Mapping + from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence from pathlib import Path from typing import Protocol from langchain.agents.middleware.human_in_the_loop import ( + ActionRequest, ApproveDecision, EditDecision, HITLRequest, @@ -29,6 +30,12 @@ from pydantic import TypeAdapter from deepagents_code._ask_user_types import AskUserWidgetResult, Question + from deepagents_code.approval_mode import ApprovalMode + from deepagents_code.hooks.client_lifecycle import ( + ClientHookContext, + ClientHookService, + ) + from deepagents_code.hooks.models.domain import DcodeNotificationKind, HookEvent from deepagents_code.resume_state import RubricResult # Type alias matching HITLResponse["decisions"] element type @@ -44,6 +51,12 @@ class _TokensShowCallback(Protocol): def __call__(self, *, approximate: bool = False) -> None: ... + class _ClientHookSessionState(Protocol): + thread_id: str + approval_mode: ApprovalMode + turn_id: str | None + client_hooks: ClientHookService | None + from deepagents_code._ask_user_types import AskUserRequest from deepagents_code._cli_context import CLIContext @@ -94,6 +107,131 @@ def __call__(self, *, approximate: bool = False) -> None: ... _ASK_USER_UNSUPPORTED_ERROR = "ask_user not supported by this UI" +class _PermissionHookOutcome(NamedTuple): + decision: dict[str, str] | None + interrupt: bool + + +def _client_hook_context(session_state: _ClientHookSessionState) -> ClientHookContext: + from deepagents_code.approval_mode import ApprovalMode, coerce_approval_mode + from deepagents_code.hooks.client_lifecycle import ClientHookContext + + return ClientHookContext.create( + thread_id=session_state.thread_id, + approval_mode=coerce_approval_mode( + getattr(session_state, "approval_mode", ApprovalMode.MANUAL) + ), + prompt_id=getattr(session_state, "turn_id", None), + ) + + +def _has_client_hook_handlers( + session_state: _ClientHookSessionState, + event: HookEvent, +) -> bool: + service = getattr(session_state, "client_hooks", None) + return service is not None and service.has_handlers(event) + + +async def _notify_client_hook( + session_state: _ClientHookSessionState, + kind: DcodeNotificationKind, + message: str, + *, + title: str | None = None, +) -> None: + from deepagents_code.hooks.client_lifecycle import ClientHookStopError + + service = getattr(session_state, "client_hooks", None) + if service is None: + return + try: + await service.notification( + _client_hook_context(session_state), + kind, + message, + title=title, + ) + except ClientHookStopError: + raise + except Exception: + logger.warning("Notification hook invocation failed", exc_info=True) + + +async def _permission_hook_outcomes( + session_state: _ClientHookSessionState, + interrupt_id: str, + action_requests: list[ActionRequest], + current_tool_messages: Mapping[str, ToolCallMessage], +) -> list[_PermissionHookOutcome]: + service = getattr(session_state, "client_hooks", None) + if service is None: + return [_PermissionHookOutcome(None, False) for _ in action_requests] + + from deepagents_code.hooks.models.domain import ToolCallData + + candidates = list(current_tool_messages.items()) + claimed: set[str] = set() + outcomes: list[_PermissionHookOutcome] = [] + context = _client_hook_context(session_state) + for index, request in enumerate(action_requests): + name = request.get("name") + args = request.get("args") + if not isinstance(name, str) or not isinstance(args, dict): + outcomes.append(_PermissionHookOutcome(None, False)) + continue + tool_id = f"{interrupt_id}:{index}" + for candidate_id, tool_message in candidates: + if candidate_id in claimed: + continue + if tool_message.tool_name == name and tool_message.args == args: + tool_id = candidate_id + claimed.add(candidate_id) + break + try: + hook_decision = await service.permission_request( + context, + ToolCallData(id=tool_id, name=name, args=args), + ) + except Exception: + logger.warning("PermissionRequest hook invocation failed", exc_info=True) + outcomes.append(_PermissionHookOutcome(None, False)) + continue + if not hook_decision.continue_processing: + reason = hook_decision.stop_reason or "Permission stopped by hook" + outcomes.append( + _PermissionHookOutcome( + {"type": "reject", "message": reason}, + True, + ) + ) + continue + permission = hook_decision.permission + if permission.behavior == "allow": + outcomes.append(_PermissionHookOutcome({"type": "approve"}, False)) + elif permission.behavior == "deny": + decision = {"type": "reject"} + if permission.reason: + decision["message"] = permission.reason + outcomes.append(_PermissionHookOutcome(decision, permission.interrupt)) + else: + outcomes.append(_PermissionHookOutcome(None, False)) + return outcomes + + +def _merge_permission_outcomes( + outcomes: list[_PermissionHookOutcome], + reviewed: Sequence[HITLDecision], +) -> list[HITLDecision]: + reviewed_iter = iter(reviewed) + return [ + cast("HITLDecision", outcome.decision) + if outcome.decision is not None + else next(reviewed_iter) + for outcome in outcomes + ] + + def _dispatch_tool_use_hook( tool_name: str, tool_id: str, tool_args: dict[str, Any] ) -> None: @@ -752,6 +890,7 @@ async def execute_task_textual( wall-clock time). Raises: + ClientHookStopError: If a compact lifecycle hook stops processing. ValidationError: If HITL request validation fails (re-raised). RuntimeError: If Manual cannot be persisted before graph execution. """ @@ -919,8 +1058,20 @@ def _notify_user_visible_output_started() -> None: turn_id=turn_id, ) user_msg["additional_kwargs"] = trusted_kwargs + messages: list[dict[str, Any]] = [] + client_hooks = getattr(session_state, "client_hooks", None) + if client_hooks is not None: + session_context = client_hooks.take_session_context(thread_id) + if session_context: + messages.append( + { + "role": "system", + "content": "\n\n".join(session_context), + } + ) + messages.append(user_msg) stream_input: dict | Command = { - "messages": [user_msg], + "messages": messages, "goal_criteria_request": None, } if rubric: @@ -933,6 +1084,7 @@ def _notify_user_visible_output_started() -> None: # Track summarization lifecycle so spinner status and notification stay in sync. summarization_in_progress = False + summarization_observed = False try: while True: @@ -1010,6 +1162,7 @@ def _notify_user_visible_output_started() -> None: from deepagents_code.hooks.client import fulfill_hook_interrupt from deepagents_code.hooks.context import apply_hooks_context from deepagents_code.hooks.interrupt import is_hook_interrupt_payload + from deepagents_code.hooks.models.domain import HookEvent apply_hooks_context( context, @@ -1148,7 +1301,9 @@ def _notify_user_visible_output_started() -> None: iv = interrupt_obj.value if is_hook_interrupt_payload(iv): hooks_runtime = getattr( - session_state, "hooks_runtime", None + session_state, + "hooks_runtime", + None, ) if hooks_runtime is None: msg = ( @@ -1237,7 +1392,11 @@ def _notify_user_visible_output_started() -> None: tool_id ] = tool_msg interrupt_occurred = True - await dispatch_hook("input.required", {}) + if not _has_client_hook_handlers( + session_state, + HookEvent.NOTIFICATION, + ): + await dispatch_hook("input.required", {}) except ValidationError: logger.exception( "Invalid ask_user interrupt payload" @@ -1253,7 +1412,11 @@ def _notify_user_visible_output_started() -> None: validated_request, ) interrupt_occurred = True - await dispatch_hook("input.required", {}) + if not _has_client_hook_handlers( + session_state, + HookEvent.NOTIFICATION, + ): + await dispatch_hook("input.required", {}) except ValidationError: # noqa: TRY203 # Re-raise preserves exception context in handler raise @@ -1294,6 +1457,7 @@ def _notify_user_visible_output_started() -> None: # These are hidden from the user; only the spinner and a # notification widget provide feedback. if _is_summarization_chunk(metadata): + summarization_observed = True if not summarization_in_progress: summarization_in_progress = True if adapter._set_spinner: @@ -1712,6 +1876,30 @@ def _notify_user_visible_output_started() -> None: ) if adapter._set_spinner and not adapter._current_tool_messages: await adapter._set_spinner("Thinking") + if summarization_observed: + from deepagents_code.hooks.client_lifecycle import ClientHookStopError + from deepagents_code.hooks.models.domain import SessionStartCause + + service = getattr(session_state, "client_hooks", None) + if service is not None: + try: + decision = await service.session_start( + _client_hook_context(session_state), + SessionStartCause.COMPACT, + ) + except Exception: + logger.warning( + "Compact SessionStart hook invocation failed", + exc_info=True, + ) + else: + if not decision.continue_processing: + reason = ( + decision.stop_reason + or "Compact session start stopped by hook" + ) + raise ClientHookStopError(reason) + summarization_observed = False # Flush any remaining text from all namespaces for ns_key, pending_text in list(pending_text_by_namespace.items()): @@ -1774,6 +1962,15 @@ def _notify_user_visible_output_started() -> None: tool_args = {"questions": questions} if adapter._request_ask_user: + from deepagents_code.hooks.models.domain import ( + DcodeNotificationKind, + ) + + await _notify_client_hook( + session_state, + DcodeNotificationKind.AGENT_NEEDS_INPUT, + "Agent needs input", + ) if adapter._set_spinner: await adapter._set_spinner(None) result: AskUserWidgetResult | dict[str, str] = { @@ -1969,9 +2166,11 @@ def _notify_user_visible_output_started() -> None: ): action_requests = hitl_request["action_requests"] - if ( - getattr(session_state, "approval_mode", None) - is ApprovalMode.YOLO + if getattr( + session_state, "approval_mode", None + ) is ApprovalMode.YOLO and not _has_client_hook_handlers( + session_state, + HookEvent.PERMISSION_REQUEST, ): decisions: list[HITLDecision] = [ ApproveDecision(type="approve") for _ in action_requests @@ -1985,6 +2184,124 @@ def _notify_user_visible_output_started() -> None: tool_msg.set_running() adapter._sync_tool_widget(tool_msg) else: + all_action_requests = action_requests + hook_outcomes = await _permission_hook_outcomes( + session_state, + interrupt_id, + all_action_requests, + adapter._current_tool_messages, + ) + if any(outcome.interrupt for outcome in hook_outcomes): + decisions = [ + cast( + "HITLDecision", + outcome.decision + or { + "type": "reject", + "message": "Permission interrupted by hook", + }, + ) + for outcome in hook_outcomes + ] + for tool_msg in _interrupt_tool_rows( + namespace, + all_action_requests, + adapter._current_tool_messages, + ): + tool_msg.set_rejected(reason="Permission interrupted") + adapter._sync_tool_widget(tool_msg) + resume_payload[interrupt_id] = {"decisions": decisions} + any_rejected = True + break + + action_requests = [ + request + for request, outcome in zip( + all_action_requests, + hook_outcomes, + strict=True, + ) + if outcome.decision is None + ] + resolved_row_ids: set[int] = set() + for request, outcome in zip( + all_action_requests, + hook_outcomes, + strict=True, + ): + if outcome.decision is None: + continue + rows = _interrupt_owned_tool_rows( + [request], + adapter._current_tool_messages, + ) + for tool_msg in rows: + resolved_row_ids.add(id(tool_msg)) + if outcome.decision["type"] == "approve": + tool_msg.set_running() + tool_name = request.get("name") + args = request.get("args") + if tool_name in { + "write_file", + "edit_file", + "delete", + } and isinstance(args, dict): + file_op_tracker.mark_hitl_approved( + tool_name, + args, + ) + else: + tool_msg.set_rejected( + reason=outcome.decision.get("message") + ) + adapter._sync_tool_widget(tool_msg) + + if not action_requests: + decisions = _merge_permission_outcomes(hook_outcomes, []) + for tool_msg in adapter._current_tool_messages.values(): + if id(tool_msg) not in resolved_row_ids: + tool_msg.set_running() + adapter._sync_tool_widget(tool_msg) + resume_payload[interrupt_id] = {"decisions": decisions} + continue + + if ( + getattr(session_state, "approval_mode", None) + is ApprovalMode.YOLO + ): + reviewed = [ + ApproveDecision(type="approve") for _ in action_requests + ] + decisions = _merge_permission_outcomes( + hook_outcomes, + reviewed, + ) + resume_payload[interrupt_id] = {"decisions": decisions} + for tool_msg in _interrupt_tool_rows( + namespace, + action_requests, + adapter._current_tool_messages, + ): + if id(tool_msg) in resolved_row_ids: + continue + tool_msg.set_running() + adapter._sync_tool_widget(tool_msg) + continue + + review_namespace = ( + namespace + if len(action_requests) == len(all_action_requests) + else ("permission_hook",) + ) + from deepagents_code.hooks.models.domain import ( + DcodeNotificationKind, + ) + + await _notify_client_hook( + session_state, + DcodeNotificationKind.PERMISSION_REQUIRED, + "Permission required", + ) # Batch approval - one dialog for all parallel tool calls await dispatch_hook( "permission.request", @@ -2056,7 +2373,7 @@ def _notify_user_visible_output_started() -> None: for _ in action_requests ] tool_msgs = _interrupt_tool_rows( - namespace, + review_namespace, action_requests, adapter._current_tool_messages, ) @@ -2100,7 +2417,7 @@ def _notify_user_visible_output_started() -> None: for _ in action_requests ] tool_msgs = _interrupt_tool_rows( - namespace, + review_namespace, action_requests, adapter._current_tool_messages, ) @@ -2227,6 +2544,10 @@ def _notify_user_visible_output_started() -> None: adapter._current_tool_messages.clear() any_rejected = True + decisions = _merge_permission_outcomes( + hook_outcomes, + decisions, + ) resume_payload[interrupt_id] = {"decisions": decisions} if any_rejected: @@ -2283,7 +2604,20 @@ def _notify_user_visible_output_started() -> None: # fires on cancel and mid-stream error too (not only this clean # end) — mirroring the headless surface, whose identical # diagnostic lives in `_run_agent_loop`'s `finally`. - await dispatch_hook("task.complete", {"thread_id": thread_id}) + from deepagents_code.hooks.models.domain import ( + DcodeNotificationKind, + ) + + await _notify_client_hook( + session_state, + DcodeNotificationKind.AGENT_COMPLETED, + "Agent completed", + ) + if not _has_client_hook_handlers( + session_state, + HookEvent.NOTIFICATION, + ): + await dispatch_hook("task.complete", {"thread_id": thread_id}) break except (asyncio.CancelledError, KeyboardInterrupt): diff --git a/libs/code/tests/unit_tests/hooks/test_client_lifecycle.py b/libs/code/tests/unit_tests/hooks/test_client_lifecycle.py new file mode 100644 index 0000000000..25eb128f60 --- /dev/null +++ b/libs/code/tests/unit_tests/hooks/test_client_lifecycle.py @@ -0,0 +1,320 @@ +"""Unit tests for client-owned Hooks v2 lifecycle integration.""" + +from __future__ import annotations + +from collections import deque +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Literal + +import pytest +from rich.console import Console + +from deepagents_code.approval_mode import ApprovalMode +from deepagents_code.client.non_interactive import ( + StreamState, + _process_hitl_interrupts, +) +from deepagents_code.hooks.client_lifecycle import ( + ClientHookContext, + ClientHookService, + ClientHookStopError, +) +from deepagents_code.hooks.models.domain import ( + DcodeNotificationKind, + HookDecision, + HookDiagnostic, + HookEvent, + HookInvocation, + NotificationDecision, + PermissionEffect, + PermissionRequestDecision, + SessionEndCause, + SessionEndDecision, + SessionStartCause, + SessionStartDecision, +) +from deepagents_code.hooks.models.wire import NotificationWireInput +from deepagents_code.hooks.projection import project_hook_input +from deepagents_code.tui.textual_adapter import ( + _merge_permission_outcomes, + _permission_hook_outcomes, +) + +if TYPE_CHECKING: + from pathlib import Path + + from langchain.agents.middleware.human_in_the_loop import ( + ApproveDecision, + EditDecision, + RejectDecision, + ) + + HITLDecision = ApproveDecision | EditDecision | RejectDecision + + +@dataclass(slots=True) +class _Runtime: + cwd: Path + decisions: deque[HookDecision] + invocations: list[HookInvocation] = field(default_factory=list) + + def configured_events(self) -> frozenset[HookEvent]: + return frozenset(decision.event for decision in self.decisions) + + async def invoke(self, invocation: HookInvocation) -> HookDecision: + self.invocations.append(invocation) + return self.decisions.popleft() + + +@dataclass(slots=True) +class _SessionState: + thread_id: str + approval_mode: ApprovalMode + turn_id: str | None + client_hooks: ClientHookService | None + + +def _context() -> ClientHookContext: + return ClientHookContext.create( + thread_id="thread-1", + approval_mode=ApprovalMode.MANUAL, + ) + + +def _permission( + behavior: Literal["allow", "deny", "ask", "none"], + *, + reason: str | None = None, + interrupt: bool = False, +) -> PermissionRequestDecision: + return PermissionRequestDecision( + event=HookEvent.PERMISSION_REQUEST, + permission=PermissionEffect( + behavior=behavior, + reason=reason, + interrupt=interrupt, + ), + ) + + +async def test_service_applies_common_effects_and_session_context( + tmp_path: Path, + caplog: pytest.LogCaptureFixture, + capsys: pytest.CaptureFixture[str], +) -> None: + notices: list[str] = [] + runtime = _Runtime( + cwd=tmp_path, + decisions=deque( + [ + SessionStartDecision( + event=HookEvent.SESSION_START, + context=["hook context"], + user_notices=["visible notice"], + terminal_sequences=["\a"], + diagnostics=[ + HookDiagnostic( + code="test_warning", + severity="warning", + message="diagnostic", + ) + ], + ) + ] + ), + ) + service = ClientHookService(runtime, notice=notices.append) + + decision = await service.session_start(_context(), SessionStartCause.STARTUP) + + assert decision.context == ["hook context"] + assert notices == ["visible notice"] + assert capsys.readouterr().out == "\a" + assert "test_warning" in caplog.text + assert service.take_session_context("thread-1") == ("hook context",) + assert service.take_session_context("thread-1") == () + + +@pytest.mark.parametrize( + ("kind", "wire_type"), + [ + (DcodeNotificationKind.PERMISSION_REQUIRED, "permission_prompt"), + (DcodeNotificationKind.AGENT_NEEDS_INPUT, "agent_needs_input"), + (DcodeNotificationKind.AGENT_COMPLETED, "agent_completed"), + ], +) +async def test_notification_service_maps_supported_kinds( + tmp_path: Path, + kind: DcodeNotificationKind, + wire_type: str, +) -> None: + runtime = _Runtime( + cwd=tmp_path, + decisions=deque([NotificationDecision(event=HookEvent.NOTIFICATION)]), + ) + service = ClientHookService(runtime) + + await service.notification(_context(), kind, "message") + + wire = project_hook_input( + runtime.invocations[0], + transcript_path=tmp_path / "transcript.jsonl", + ) + assert isinstance(wire, NotificationWireInput) + assert wire.notification_type == wire_type + + +async def test_session_end_discards_pending_context(tmp_path: Path) -> None: + runtime = _Runtime( + cwd=tmp_path, + decisions=deque( + [ + SessionStartDecision( + event=HookEvent.SESSION_START, + context=["pending"], + ), + SessionEndDecision(event=HookEvent.SESSION_END), + ] + ), + ) + service = ClientHookService(runtime) + context = _context() + + await service.session_start(context, SessionStartCause.RESUME) + await service.session_end(context, SessionEndCause.RESUME) + + assert service.take_session_context("thread-1") == () + assert [invocation.event.event for invocation in runtime.invocations] == [ + HookEvent.SESSION_START, + HookEvent.SESSION_END, + ] + + +async def test_notification_stop_interrupts_client_processing(tmp_path: Path) -> None: + runtime = _Runtime( + cwd=tmp_path, + decisions=deque( + [ + NotificationDecision( + event=HookEvent.NOTIFICATION, + continue_processing=False, + stop_reason="stop now", + ) + ] + ), + ) + service = ClientHookService(runtime) + + with pytest.raises(ClientHookStopError, match="stop now"): + await service.notification( + _context(), + DcodeNotificationKind.AGENT_COMPLETED, + "done", + ) + + +@pytest.mark.parametrize( + ("behavior", "hook_decision", "reviewed", "expected"), + [ + ("allow", {"type": "approve"}, [], {"type": "approve"}), + ( + "deny", + {"type": "reject", "message": "blocked"}, + [], + {"type": "reject", "message": "blocked"}, + ), + ("none", None, [{"type": "approve"}], {"type": "approve"}), + ], +) +async def test_tui_permission_decisions_precede_review( + tmp_path: Path, + behavior: Literal["allow", "deny", "none"], + hook_decision: dict[str, str] | None, + reviewed: list[HITLDecision], + expected: dict[str, str], +) -> None: + runtime = _Runtime( + cwd=tmp_path, + decisions=deque( + [ + _permission( + behavior, + reason="blocked" if behavior == "deny" else None, + ) + ] + ), + ) + state = _SessionState( + thread_id="thread-1", + approval_mode=ApprovalMode.MANUAL, + turn_id=None, + client_hooks=ClientHookService(runtime), + ) + + outcomes = await _permission_hook_outcomes( + state, + "interrupt-1", + [{"name": "read_file", "args": {"path": "README.md"}}], + {}, + ) + + assert outcomes[0].decision == hook_decision + assert _merge_permission_outcomes(outcomes, reviewed) == [expected] + assert runtime.invocations[0].event.event is HookEvent.PERMISSION_REQUEST + + +@pytest.mark.parametrize( + ("decisions", "expected"), + [ + ( + deque([_permission("allow")]), + {"type": "approve"}, + ), + ( + deque([_permission("deny", reason="blocked")]), + {"type": "reject", "message": "blocked"}, + ), + ( + deque( + [ + _permission("none"), + NotificationDecision(event=HookEvent.NOTIFICATION), + ] + ), + {"type": "approve"}, + ), + ], +) +async def test_headless_permission_decisions_precede_resolution( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + decisions: deque[HookDecision], + expected: dict[str, str], +) -> None: + runtime = _Runtime(cwd=tmp_path, decisions=decisions) + state = StreamState(client_hooks=ClientHookService(runtime)) + state.pending_interrupts["interrupt-1"] = { + "action_requests": [ + {"name": "read_file", "args": {"path": "README.md"}}, + ], + "review_configs": [], + } + resolution_calls = 0 + + def _resolve(*_args: object) -> dict[str, str]: + nonlocal resolution_calls + resolution_calls += 1 + assert runtime.invocations[-1].event.event is HookEvent.NOTIFICATION + return {"type": "approve"} + + monkeypatch.setattr( + "deepagents_code.client.non_interactive._make_hitl_decision", + _resolve, + ) + + await _process_hitl_interrupts(state, Console(quiet=True), "thread-1") + + assert state.hitl_response["interrupt-1"]["decisions"] == [expected] + should_resolve = expected == {"type": "approve"} and len(runtime.invocations) == 2 + assert resolution_calls == int(should_resolve) + assert runtime.invocations[0].event.event is HookEvent.PERMISSION_REQUEST diff --git a/libs/code/tests/unit_tests/hooks/test_engine.py b/libs/code/tests/unit_tests/hooks/test_engine.py index 9eaa074919..b1d127d2e9 100644 --- a/libs/code/tests/unit_tests/hooks/test_engine.py +++ b/libs/code/tests/unit_tests/hooks/test_engine.py @@ -22,6 +22,7 @@ AgentIdentity, CompactTrigger, DcodeNotification, + DcodeNotificationKind, HookContext, HookDiagnostic, HookEvent, @@ -189,7 +190,7 @@ def test_snapshot_matches_notification_and_skips_tool_mismatch(tmp_path: Path) - NotificationEvent( event=HookEvent.NOTIFICATION, notification=DcodeNotification( - type="permission_prompt", + type=DcodeNotificationKind.PERMISSION_REQUIRED, message="Approve", ), ), @@ -390,7 +391,7 @@ def test_snapshot_rejects_matcher_for_unmatchable_event() -> None: NotificationEvent( event=HookEvent.NOTIFICATION, notification=DcodeNotification( - type="permission_prompt", + type=DcodeNotificationKind.PERMISSION_REQUIRED, message="Approve", ), ), diff --git a/libs/code/tests/unit_tests/test_app.py b/libs/code/tests/unit_tests/test_app.py index 6f4efb2b9e..7a490ba2d8 100644 --- a/libs/code/tests/unit_tests/test_app.py +++ b/libs/code/tests/unit_tests/test_app.py @@ -555,6 +555,49 @@ async def capture_history( # noqa: RUF029 assert call_count == 1 assert app._initial_session_started is True + async def test_session_start_waits_for_session_runtime(self) -> None: + """Mounted startup cannot outrun asynchronous session initialization.""" + app = DeepAgentsApp( + agent=MagicMock(), + thread_id="thread-123", + initial_prompt="hello", + ) + run_hook = AsyncMock(return_value=True) + submit = AsyncMock() + app._run_session_start_hook = run_hook # ty: ignore[invalid-assignment] + app._submit_initial_submission = submit # ty: ignore[invalid-assignment] + app._session_init_started = True + + task = asyncio.create_task(app._run_session_start_sequence()) + await asyncio.sleep(0) + run_hook.assert_not_awaited() + submit.assert_not_awaited() + + app._session_state = TextualSessionState(thread_id="thread-123") + app._session_state_ready.set() + await task + + run_hook.assert_awaited_once() + submit.assert_awaited_once() + + async def test_stopped_session_start_blocks_initial_submission(self) -> None: + """A hook stop at startup prevents the initial prompt from running.""" + app = DeepAgentsApp( + agent=MagicMock(), + thread_id="thread-123", + initial_prompt="hello", + ) + app._session_state = TextualSessionState(thread_id="thread-123") + app._run_session_start_hook = AsyncMock( # ty: ignore[invalid-assignment] + return_value=False + ) + submit = AsyncMock() + app._submit_initial_submission = submit # ty: ignore[invalid-assignment] + + await app._run_session_start_sequence() + + submit.assert_not_awaited() + async def test_reconnect_drains_queue_without_reloading_history(self) -> None: """Later `ServerReady` events should drain queued input once connected.""" app = DeepAgentsApp( diff --git a/libs/code/tests/unit_tests/test_non_interactive.py b/libs/code/tests/unit_tests/test_non_interactive.py index 3a51666334..67eed2381d 100644 --- a/libs/code/tests/unit_tests/test_non_interactive.py +++ b/libs/code/tests/unit_tests/test_non_interactive.py @@ -24,6 +24,7 @@ UNRENDERABLE_TOOL_OUTPUT, ToolCallBuffer, ) +from deepagents_code.approval_mode import ApprovalMode from deepagents_code.client.non_interactive import ( _MAX_HITL_ITERATIONS, HITLIterationLimitError, @@ -44,6 +45,7 @@ ) from deepagents_code.config import SHELL_ALLOW_ALL, ModelResult from deepagents_code.file_ops import FileOpTracker +from deepagents_code.hooks.models.domain import HookEvent from deepagents_code.tool_display import format_tool_message_content @@ -321,6 +323,61 @@ async def test_sandbox_type_passed_to_server(self) -> None: assert kwargs["profile_overrides"] == {"max_input_tokens": 32_000} assert kwargs["enable_interpreter"] is None + async def test_permission_hooks_override_headless_yolo_bypass(self) -> None: + """Permission hooks force client resolution while retaining YOLO context.""" + runtime = MagicMock() + runtime.configured_events.return_value = frozenset( + {HookEvent.PERMISSION_REQUEST} + ) + mock_agent = MagicMock() + mock_server_proc = MagicMock() + + with ( + patch( + "deepagents_code.client.non_interactive.create_model", + return_value=ModelResult( + model=MagicMock(), + model_name="test-model", + provider="test", + ), + ), + patch( + "deepagents_code.client.non_interactive.generate_thread_id", + return_value="test-thread", + ), + patch("deepagents_code.client.non_interactive.settings") as mock_settings, + patch( + "deepagents_code.client.non_interactive.build_langsmith_thread_url", + return_value=None, + ), + patch( + "deepagents_code.hooks.runtime.HooksRuntime.create", + return_value=runtime, + ), + patch( + "deepagents_code.client.non_interactive._run_agent_loop", + new_callable=AsyncMock, + ) as mock_loop, + patch( + "deepagents_code.client.launch.server_manager.start_server_and_get_agent", + new_callable=AsyncMock, + return_value=(mock_agent, mock_server_proc, None), + ) as mock_start_server, + ): + mock_settings.shell_allow_list = SHELL_ALLOW_ALL + mock_settings.has_tavily = False + mock_settings.model_name = None + + await run_non_interactive(message="test task") + + _, server_kwargs = mock_start_server.call_args + assert server_kwargs["auto_approve"] is False + assert server_kwargs["interrupt_shell_only"] is False + _, loop_kwargs = mock_loop.call_args + assert loop_kwargs["hooks_runtime"] is runtime + assert loop_kwargs["approval_mode"] is ApprovalMode.YOLO + assert loop_kwargs["prompt_id"] is not None + async def test_sandbox_snapshot_name_passed_to_server(self) -> None: """`sandbox_snapshot_name` must reach `start_server_and_get_agent`.""" mock_agent = MagicMock() diff --git a/libs/code/tests/unit_tests/tui/test_textual_adapter.py b/libs/code/tests/unit_tests/tui/test_textual_adapter.py index 614d591958..e92d11255a 100644 --- a/libs/code/tests/unit_tests/tui/test_textual_adapter.py +++ b/libs/code/tests/unit_tests/tui/test_textual_adapter.py @@ -35,6 +35,11 @@ _process_message_chunk, ) from deepagents_code.config import ASCII_GLYPHS, UNICODE_GLYPHS, build_stream_config +from deepagents_code.hooks.models.domain import ( + HookEvent, + PermissionEffect, + PermissionRequestDecision, +) from deepagents_code.tui.textual_adapter import ( RubricEvaluationEnd, SessionStats, @@ -6265,6 +6270,65 @@ async def request_approval( "tool_output": "Tool approval rejected", } + async def test_yolo_permission_hook_can_reject_before_auto_approval(self) -> None: + """YOLO still invokes configured permission hooks before resolution.""" + action_requests = [{"name": "execute", "args": {"command": "echo hi"}}] + agent = _SequencedAgent( + streams_by_call=[ + [ + _hitl_interrupt_chunk( + { + "action_requests": action_requests, + "review_configs": [], + } + ) + ], + [], + ] + ) + client_hooks = MagicMock() + client_hooks.has_handlers.side_effect = lambda event: ( + event is HookEvent.PERMISSION_REQUEST + ) + client_hooks.permission_request = AsyncMock( + return_value=PermissionRequestDecision( + event=HookEvent.PERMISSION_REQUEST, + permission=PermissionEffect( + behavior="deny", + reason="blocked by hook", + ), + ) + ) + request_approval = AsyncMock() + adapter = TextualUIAdapter( + mount_message=_mock_mount, + update_status=_noop_status, + request_approval=request_approval, + ) + + await execute_task_textual( + user_input="hello", + agent=agent, + assistant_id="assistant", + session_state=SimpleNamespace( + thread_id="thread-1", + approval_mode=ApprovalMode.YOLO, + auto_approve=True, + turn_id=None, + client_hooks=client_hooks, + ), + adapter=adapter, + ) + + request_approval.assert_not_awaited() + client_hooks.permission_request.assert_awaited_once() + resume_cmd = agent.stream_inputs[1] + assert isinstance(resume_cmd, Command) + resume_payload = cast("dict[str, dict[str, Any]]", resume_cmd.resume) + assert resume_payload["interrupt-1"]["decisions"] == [ + {"type": "reject", "message": "blocked by hook"} + ] + async def test_hitl_reasoned_reject_preserves_tool_args_for_result(self) -> None: """A reasoned HITL reject keeps args until the resumed ToolMessage.""" action_requests = [{"name": "execute", "args": {"command": "echo hi"}}] From 2d32caff745f018b27c5708d7ee641b7eba517dc Mon Sep 17 00:00:00 2001 From: Johannes du Plessis Date: Fri, 24 Jul 2026 08:53:05 -0700 Subject: [PATCH 09/14] fix(code): avoid session-start deadlock on hooks init Initialize session state inline before SessionStart instead of waiting on a Textual worker Event that never progresses under that wait. Co-authored-by: Cursor --- libs/code/deepagents_code/app.py | 96 ++++++++++++++++++-------------- 1 file changed, 53 insertions(+), 43 deletions(-) diff --git a/libs/code/deepagents_code/app.py b/libs/code/deepagents_code/app.py index 706587b27d..8213de7ccd 100644 --- a/libs/code/deepagents_code/app.py +++ b/libs/code/deepagents_code/app.py @@ -3557,6 +3557,7 @@ def __init__( """ self._session_state_ready = asyncio.Event() self._session_init_started = False + self._session_init_lock = asyncio.Lock() self._startup_task: asyncio.Task[None] | None = None """Startup task reference (set in on_mount).""" @@ -4289,53 +4290,59 @@ async def _post_paint_init(self) -> None: async def _init_session_state(self) -> None: """Create session state in a thread (imports deepagents_code.sessions).""" - self._session_init_started = True + async with self._session_init_lock: + if self._session_state is not None: + self._session_state_ready.set() + return + self._session_init_started = True - def _create() -> TextualSessionState: - from pathlib import Path + def _create() -> TextualSessionState: + from pathlib import Path - from deepagents_code.hooks.runtime import HooksRuntime + from deepagents_code.hooks.runtime import HooksRuntime - state = TextualSessionState( - approval_mode=self._approval_mode, - thread_id=self._lc_thread_id, - ) - try: - # Interactive sessions keep project hooks off until a dedicated - # workspace-trust prompt lands (design-doc security follow-up). - state.hooks_runtime = HooksRuntime.create( - cwd=Path(self._cwd), - workspace_trusted=False, + state = TextualSessionState( + approval_mode=self._approval_mode, + thread_id=self._lc_thread_id, ) - except Exception: - logger.exception("Failed to create HooksRuntime; server hooks disabled") - state.hooks_runtime = None - return state + try: + # Interactive sessions keep project hooks off until a dedicated + # workspace-trust prompt lands (design-doc security follow-up). + state.hooks_runtime = HooksRuntime.create( + cwd=Path(self._cwd), + workspace_trusted=False, + ) + except Exception: + logger.exception( + "Failed to create HooksRuntime; server hooks disabled" + ) + state.hooks_runtime = None + return state - try: - session_state = await asyncio.to_thread(_create) - except Exception: - logger.exception("Failed to create session state") - self.notify( - "Session initialization failed. Some features may be unavailable.", - severity="error", - timeout=10, - ) + try: + session_state = await asyncio.to_thread(_create) + except Exception: + logger.exception("Failed to create session state") + self.notify( + "Session initialization failed. Some features may be unavailable.", + severity="error", + timeout=10, + ) + self._session_state_ready.set() + return + # A user can change the approval mode while session construction runs + # in the worker thread. Re-read the app-owned selection on the event + # loop so the newly assigned state cannot overwrite that newer choice. + session_state.approval_mode = self._approval_mode + if session_state.hooks_runtime is not None: + from deepagents_code.hooks.client_lifecycle import ClientHookService + + session_state.client_hooks = ClientHookService( + session_state.hooks_runtime, + notice=lambda message: self.notify(message, markup=False), + ) + self._session_state = session_state self._session_state_ready.set() - return - # A user can change the approval mode while session construction runs - # in the worker thread. Re-read the app-owned selection on the event - # loop so the newly assigned state cannot overwrite that newer choice. - session_state.approval_mode = self._approval_mode - if session_state.hooks_runtime is not None: - from deepagents_code.hooks.client_lifecycle import ClientHookService - - session_state.client_hooks = ClientHookService( - session_state.hooks_runtime, - notice=lambda message: self.notify(message, markup=False), - ) - self._session_state = session_state - self._session_state_ready.set() await self._auto_accept_pending_goal_rubric() async def _refresh_client_hooks_runtime(self) -> None: @@ -8196,8 +8203,11 @@ async def _run_session_start_sequence(self) -> None: self._schedule_session_start_after_launch_init(launch_init_task) return - if self._session_state is None and self._session_init_started: - await self._session_state_ready.wait() + if self._session_state is None: + # Initialize inline. Waiting on `_session_state_ready` while the + # session-init Textual worker is in-flight deadlocks under Textual's + # worker scheduling (the worker never progresses to set the Event). + await self._init_session_state() self._initial_session_started = True self._startup_sequence_running = True From 7b4a4ebea1676400a40280890cd5df1a23827bf7 Mon Sep 17 00:00:00 2001 From: Johannes du Plessis Date: Fri, 24 Jul 2026 09:01:42 -0700 Subject: [PATCH 10/14] fix(code): keep hook transcripts in the global config dir Stop deriving the transcript store from `config_dir`, which let tests and project-local overrides create `.deepagents/transcripts` under the package tree. Transcripts now always default to `~/.deepagents/transcripts`. Co-authored-by: Cursor --- libs/code/deepagents_code/hooks/runtime.py | 14 +++++--- .../unit_tests/hooks/test_server_lifecycle.py | 12 +++++-- .../tests/unit_tests/hooks/test_transcript.py | 33 +++++++++++++++++-- 3 files changed, 49 insertions(+), 10 deletions(-) diff --git a/libs/code/deepagents_code/hooks/runtime.py b/libs/code/deepagents_code/hooks/runtime.py index e221983fab..c8fb02be29 100644 --- a/libs/code/deepagents_code/hooks/runtime.py +++ b/libs/code/deepagents_code/hooks/runtime.py @@ -67,9 +67,10 @@ def create( cwd: Session working directory. workspace_trusted: Whether project-scoped hooks may be loaded. config_dir: Alternate user config directory for tests. - transcript_root: Alternate transcript store root. Defaults to - `~/.deepagents/transcripts`, or `{config_dir}/transcripts` when - an alternate user configuration directory is provided. + transcript_root: Alternate transcript store root for tests. + Defaults to `~/.deepagents/transcripts` regardless of + `config_dir` (project and test hook configs must not relocate + the global transcript store). Returns: A runtime ready to execute invocations for this session. @@ -84,8 +85,11 @@ def create( diagnostics=loaded.diagnostics, snapshot_id=loaded.snapshot_id, ) - user_config_dir = config_dir or DEFAULT_CONFIG_DIR - store = TranscriptStore(transcript_root or user_config_dir / "transcripts") + store = TranscriptStore( + transcript_root + if transcript_root is not None + else DEFAULT_CONFIG_DIR / "transcripts" + ) engine = HookEngine(snapshot) return cls(snapshot=snapshot, transcripts=store, engine=engine, cwd=cwd) diff --git a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py index bc2f5da6ea..abf401b4f7 100644 --- a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py +++ b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py @@ -125,7 +125,11 @@ def test_apply_hooks_context_sets_server_events(tmp_path: Path) -> None: '{"hooks":{"PreToolUse":[{"hooks":[{"type":"command","command":"true"}]}]}}', encoding="utf-8", ) - runtime = HooksRuntime.create(cwd=tmp_path, config_dir=config_dir) + runtime = HooksRuntime.create( + cwd=tmp_path, + config_dir=config_dir, + transcript_root=tmp_path / "transcripts", + ) context: CLIContext = {} apply_hooks_context(context, runtime, prompt_id="prompt-1") @@ -333,7 +337,11 @@ async def test_fulfill_hook_invocation_runs_engine(tmp_path: Path) -> None: config_dir = tmp_path / "config" config_dir.mkdir() (config_dir / "hooks.json").write_text('{"hooks":{}}', encoding="utf-8") - runtime = HooksRuntime.create(cwd=tmp_path, config_dir=config_dir) + runtime = HooksRuntime.create( + cwd=tmp_path, + config_dir=config_dir, + transcript_root=tmp_path / "transcripts", + ) request = _request() request = request.model_copy(update={"snapshot_id": runtime.snapshot_id}) diff --git a/libs/code/tests/unit_tests/hooks/test_transcript.py b/libs/code/tests/unit_tests/hooks/test_transcript.py index be0d7420f3..7c130c3839 100644 --- a/libs/code/tests/unit_tests/hooks/test_transcript.py +++ b/libs/code/tests/unit_tests/hooks/test_transcript.py @@ -183,15 +183,38 @@ def append(index: int) -> None: assert handle.revision == concurrent.revision("thread") -def test_runtime_stores_transcripts_outside_workspace(tmp_path: Path) -> None: +def test_runtime_stores_transcripts_outside_workspace( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: workspace = tmp_path / "workspace" config_dir = tmp_path / "config" + global_dir = tmp_path / "global-deepagents" workspace.mkdir() + monkeypatch.setattr( + "deepagents_code.hooks.runtime.DEFAULT_CONFIG_DIR", + global_dir, + ) runtime = HooksRuntime.create(cwd=workspace, config_dir=config_dir) - assert runtime.transcripts.root == (config_dir / "transcripts").resolve() + assert runtime.transcripts.root == (global_dir / "transcripts").resolve() assert not (workspace / ".deepagents").exists() + assert not (config_dir / "transcripts").exists() + + +def test_runtime_accepts_explicit_transcript_root(tmp_path: Path) -> None: + workspace = tmp_path / "workspace" + transcript_root = tmp_path / "isolated-transcripts" + workspace.mkdir() + + runtime = HooksRuntime.create( + cwd=workspace, + config_dir=tmp_path / "config", + transcript_root=transcript_root, + ) + + assert runtime.transcripts.root == transcript_root.resolve() + assert not (tmp_path / "config" / "transcripts").exists() async def test_runtime_materializes_paths_and_invokes(tmp_path: Path) -> None: @@ -229,7 +252,11 @@ async def test_runtime_materializes_paths_and_invokes(tmp_path: Path) -> None: ), encoding="utf-8", ) - runtime = HooksRuntime.create(cwd=tmp_path, config_dir=config_dir) + runtime = HooksRuntime.create( + cwd=tmp_path, + config_dir=config_dir, + transcript_root=tmp_path / "transcripts", + ) runtime.append_messages("thread-1", [HumanMessage(content="hi")]) invocation = HookInvocation( context=HookContext( From 3020826293509a1f0219cb39daf3f4007886f2ae Mon Sep 17 00:00:00 2001 From: Johannes du Plessis Date: Fri, 24 Jul 2026 09:04:05 -0700 Subject: [PATCH 11/14] fix(code): initialize session hooks state synchronously Avoid `to_thread` during session start so server-ready startup completes within the same event-loop turns that tests and the UI expect. Co-authored-by: Cursor --- libs/code/deepagents_code/app.py | 12 +++++++----- libs/code/tests/unit_tests/test_app.py | 15 ++++----------- 2 files changed, 11 insertions(+), 16 deletions(-) diff --git a/libs/code/deepagents_code/app.py b/libs/code/deepagents_code/app.py index 8213de7ccd..a40bb16c6a 100644 --- a/libs/code/deepagents_code/app.py +++ b/libs/code/deepagents_code/app.py @@ -4289,7 +4289,7 @@ async def _post_paint_init(self) -> None: ) async def _init_session_state(self) -> None: - """Create session state in a thread (imports deepagents_code.sessions).""" + """Create session state (hooks runtime + client hook service).""" async with self._session_init_lock: if self._session_state is not None: self._session_state_ready.set() @@ -4320,7 +4320,10 @@ def _create() -> TextualSessionState: return state try: - session_state = await asyncio.to_thread(_create) + # Keep construction on the event loop. `HooksRuntime.create` is + # cheap (config load), and `to_thread` races server-ready startup + # tests that only yield a few event-loop turns. + session_state = _create() except Exception: logger.exception("Failed to create session state") self.notify( @@ -4330,9 +4333,8 @@ def _create() -> TextualSessionState: ) self._session_state_ready.set() return - # A user can change the approval mode while session construction runs - # in the worker thread. Re-read the app-owned selection on the event - # loop so the newly assigned state cannot overwrite that newer choice. + # Re-read the app-owned selection so a mode change during construction + # cannot be overwritten by the freshly built state. session_state.approval_mode = self._approval_mode if session_state.hooks_runtime is not None: from deepagents_code.hooks.client_lifecycle import ClientHookService diff --git a/libs/code/tests/unit_tests/test_app.py b/libs/code/tests/unit_tests/test_app.py index 7a490ba2d8..55bb362086 100644 --- a/libs/code/tests/unit_tests/test_app.py +++ b/libs/code/tests/unit_tests/test_app.py @@ -555,8 +555,8 @@ async def capture_history( # noqa: RUF029 assert call_count == 1 assert app._initial_session_started is True - async def test_session_start_waits_for_session_runtime(self) -> None: - """Mounted startup cannot outrun asynchronous session initialization.""" + async def test_session_start_initializes_session_runtime_inline(self) -> None: + """Startup builds session/hooks state before client SessionStart runs.""" app = DeepAgentsApp( agent=MagicMock(), thread_id="thread-123", @@ -566,17 +566,10 @@ async def test_session_start_waits_for_session_runtime(self) -> None: submit = AsyncMock() app._run_session_start_hook = run_hook # ty: ignore[invalid-assignment] app._submit_initial_submission = submit # ty: ignore[invalid-assignment] - app._session_init_started = True - task = asyncio.create_task(app._run_session_start_sequence()) - await asyncio.sleep(0) - run_hook.assert_not_awaited() - submit.assert_not_awaited() - - app._session_state = TextualSessionState(thread_id="thread-123") - app._session_state_ready.set() - await task + await app._run_session_start_sequence() + assert app._session_state is not None run_hook.assert_awaited_once() submit.assert_awaited_once() From 8c0c94c78f0bbd86df54d463c70aeeb9a6521526 Mon Sep 17 00:00:00 2001 From: Johannes du Plessis Date: Fri, 24 Jul 2026 10:42:28 -0700 Subject: [PATCH 12/14] fix(code): make Hooks v2 lifecycle replay-safe Preserve server hook identity across graph replay, run pre-tool policy before HITL, and deduplicate client fulfillment side effects. Co-authored-by: Cursor --- libs/code/deepagents_code/agent.py | 49 ++- libs/code/deepagents_code/auto_mode.py | 17 +- libs/code/deepagents_code/hooks/client.py | 73 +++- libs/code/deepagents_code/hooks/engine.py | 12 +- libs/code/deepagents_code/hooks/envelope.py | 78 ++++ .../deepagents_code/hooks/models/domain.py | 2 +- libs/code/deepagents_code/hooks/runtime.py | 10 +- .../hooks/server_middleware.py | 369 +++++++++++++---- .../tests/unit_tests/hooks/test_engine.py | 50 ++- .../unit_tests/hooks/test_server_lifecycle.py | 372 +++++++++++++++++- 10 files changed, 920 insertions(+), 112 deletions(-) create mode 100644 libs/code/deepagents_code/hooks/envelope.py diff --git a/libs/code/deepagents_code/agent.py b/libs/code/deepagents_code/agent.py index 8ecdc186ab..7af98f54eb 100644 --- a/libs/code/deepagents_code/agent.py +++ b/libs/code/deepagents_code/agent.py @@ -1911,6 +1911,14 @@ def _should_interrupt_tool_call( Returns: `True` to interrupt, or `False` for Auto/YOLO bypass. """ + from deepagents_code.hooks.server_middleware import pre_tool_behavior + + tool_call = getattr(request, "tool_call", None) + tool_call_id = str(tool_call.get("id") or "") if isinstance(tool_call, dict) else "" + hook_behavior = pre_tool_behavior(getattr(request, "state", None), tool_call_id) + if hook_behavior in {"allow", "deny"}: + return False + runtime = getattr(request, "runtime", None) mode = _async_routing_mode(getattr(request, "state", None)) if mode is None: @@ -2436,7 +2444,13 @@ def _subagent_cli_middleware( from deepagents_code.hooks.server_middleware import ServerHooksMiddleware hooks_cwd = Path(effective_cwd) if effective_cwd is not None else Path.cwd() - middleware.append(ServerHooksMiddleware(cwd=hooks_cwd, emit_stop=False)) + middleware.append( + ServerHooksMiddleware( + cwd=hooks_cwd, + emit_stop=False, + mcp_tools=mcp_tools, + ) + ) # Subagents share the on-disk filesystem backend and can edit the user # AGENTS.md, so they get the same managed onboarding-name block guard as # the main agent. Gated on memory because the block only exists when @@ -2727,7 +2741,9 @@ def _subagent_cli_middleware( from deepagents_code.hooks.server_middleware import ServerHooksMiddleware hooks_cwd = Path(effective_cwd) if effective_cwd is not None else Path.cwd() - agent_middleware.append(ServerHooksMiddleware(cwd=hooks_cwd)) + if resolved_interrupt_on is not None: + agent_middleware.append(AsyncApprovalHITLMiddleware(resolved_interrupt_on)) + agent_middleware.append(ServerHooksMiddleware(cwd=hooks_cwd, mcp_tools=mcp_tools)) # Get or use custom system prompt if system_prompt is None: @@ -2739,24 +2755,19 @@ def _subagent_cli_middleware( fs_tools=fs_tools, ) - interrupt_on: dict[str, bool | InterruptOnConfig] | None + interrupt_on: dict[str, bool | InterruptOnConfig] = {} auto_mode_config: tuple[Path, list[str]] | None = None - if resolved_interrupt_on is None: - interrupt_on = {} - else: - interrupt_on = resolved_interrupt_on # ty: ignore[invalid-assignment] # InterruptOnConfig is compatible at runtime - if auto_mode_enabled: - configured_allow_list = shell_allow_list or settings.shell_allow_list - narrow_allow_list = ( - configured_allow_list if isinstance(configured_allow_list, list) else [] - ) - trusted_root = ( - project_context.project_root - if project_context is not None - and project_context.project_root is not None - else effective_cwd or Path.cwd() - ) - auto_mode_config = (Path(trusted_root), narrow_allow_list) + if resolved_interrupt_on is not None and auto_mode_enabled: + configured_allow_list = shell_allow_list or settings.shell_allow_list + narrow_allow_list = ( + configured_allow_list if isinstance(configured_allow_list, list) else [] + ) + trusted_root = ( + project_context.project_root + if project_context is not None and project_context.project_root is not None + else effective_cwd or Path.cwd() + ) + auto_mode_config = (Path(trusted_root), narrow_allow_list) # Set up composite backend with routing. if sandbox is None: diff --git a/libs/code/deepagents_code/auto_mode.py b/libs/code/deepagents_code/auto_mode.py index 04f032b660..1e7f70a756 100644 --- a/libs/code/deepagents_code/auto_mode.py +++ b/libs/code/deepagents_code/auto_mode.py @@ -2568,6 +2568,13 @@ async def aafter_model( ) if ai_message is None or not ai_message.tool_calls: return {"_auto_decision_plan": None} + from deepagents_code.hooks.server_middleware import pre_tool_behavior + + hook_bypass_ids = { + _tool_call_id(call) + for call in ai_message.tool_calls + if pre_tool_behavior(state, _tool_call_id(call)) in {"allow", "deny"} + } thread_key = _thread_key(runtime) plan = self._validated_plan(state, ai_message, thread_key) current_mode, current_mode_unavailable = await _live_mode(runtime) @@ -2575,6 +2582,7 @@ async def aafter_model( _tool_call_id(call) for call in ai_message.tool_calls if call["name"] in self.interrupt_on + and _tool_call_id(call) not in hook_bypass_ids } if plan is None: if not manual_ids: @@ -2624,6 +2632,9 @@ async def aafter_model( current_mode = ApprovalMode.MANUAL if proposal_mode is ApprovalMode.MANUAL or current_mode is ApprovalMode.MANUAL: + review_ids = set(plan["manual_gated_ids"]) - hook_bypass_ids + if not review_ids: + return {"_auto_decision_plan": None} manual_fallback = plan["fallback_reason"] in { "approval_mode_unavailable", "control_state_unavailable", @@ -2637,7 +2648,7 @@ async def aafter_model( state, runtime, ai_message, - set(plan["manual_gated_ids"]), + review_ids, fallback=manual_fallback, counters=counters, all_manual_ids=manual_ids, @@ -2652,7 +2663,9 @@ async def aafter_model( return {"_auto_decision_plan": None} decision_by_id = { - decision["tool_call_id"]: decision for decision in plan["decisions"] + decision["tool_call_id"]: decision + for decision in plan["decisions"] + if decision["tool_call_id"] not in hook_bypass_ids } human_ids = { tool_id diff --git a/libs/code/deepagents_code/hooks/client.py b/libs/code/deepagents_code/hooks/client.py index dae168b435..f4997c0bdd 100644 --- a/libs/code/deepagents_code/hooks/client.py +++ b/libs/code/deepagents_code/hooks/client.py @@ -2,9 +2,12 @@ from __future__ import annotations +import asyncio import logging import sys +from dataclasses import dataclass, field from typing import TYPE_CHECKING +from uuid import UUID from deepagents_code.hooks.interrupt import ( build_hook_resume_value, @@ -13,7 +16,7 @@ from deepagents_code.hooks.models.transport import HookInvocationResponse if TYPE_CHECKING: - from collections.abc import Mapping + from collections.abc import Awaitable, Callable, Mapping from deepagents_code.hooks.models.domain import HookDecision from deepagents_code.hooks.models.transport import HookInvocationRequest @@ -21,6 +24,54 @@ logger = logging.getLogger(__name__) +_FulfillmentKey = tuple[str, UUID] +_ResumePayload = dict[str, object] + + +@dataclass(slots=True) +class HookFulfillmentLedger: + """Deduplicate hook fulfillment for one client session.""" + + _in_flight: dict[_FulfillmentKey, asyncio.Task[HookInvocationResponse]] = field( + default_factory=dict + ) + _completed: dict[_FulfillmentKey, HookInvocationResponse] = field( + default_factory=dict + ) + _lock: asyncio.Lock = field(default_factory=asyncio.Lock) + + async def fulfill( + self, + key: _FulfillmentKey, + operation: Callable[[], Awaitable[HookInvocationResponse]], + ) -> HookInvocationResponse: + """Return one shared result for concurrent and repeated delivery.""" + async with self._lock: + completed = self._completed.get(key) + if completed is not None: + return completed + task = self._in_flight.get(key) + if task is None: + task = asyncio.create_task(self._run(key, operation)) + self._in_flight[key] = task + return await asyncio.shield(task) + + async def _run( + self, + key: _FulfillmentKey, + operation: Callable[[], Awaitable[HookInvocationResponse]], + ) -> HookInvocationResponse: + try: + result = await operation() + except BaseException: + async with self._lock: + self._in_flight.pop(key, None) + raise + async with self._lock: + self._completed[key] = result + self._in_flight.pop(key, None) + return result + async def fulfill_hook_invocation( runtime: HooksRuntime, @@ -45,13 +96,19 @@ async def fulfill_hook_invocation( ) raise ValueError(msg) - decision = await runtime.invoke(request.invocation) - _apply_client_side_effects(decision) - response = HookInvocationResponse( - protocol_version=1, - invocation_id=request.invocation_id, - snapshot_id=request.snapshot_id, - decision=decision, + async def execute() -> HookInvocationResponse: + decision = await runtime.invoke(request.invocation) + _apply_client_side_effects(decision) + return HookInvocationResponse( + protocol_version=1, + invocation_id=request.invocation_id, + snapshot_id=request.snapshot_id, + decision=decision, + ) + + response = await runtime.fulfillments.fulfill( + (request.snapshot_id, request.invocation_id), + execute, ) return build_hook_resume_value(response) diff --git a/libs/code/deepagents_code/hooks/engine.py b/libs/code/deepagents_code/hooks/engine.py index 98e513a8fc..9ba35b2234 100644 --- a/libs/code/deepagents_code/hooks/engine.py +++ b/libs/code/deepagents_code/hooks/engine.py @@ -3,13 +3,12 @@ from __future__ import annotations import asyncio -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import TYPE_CHECKING from deepagents_code.hooks.capabilities import get_event_spec +from deepagents_code.hooks.envelope import HookEnvelopeAdapter from deepagents_code.hooks.models.domain import HookDiagnostic -from deepagents_code.hooks.projection import serialize_hook_input -from deepagents_code.hooks.reducer import reduce_hook_results from deepagents_code.hooks.runner import ( MAX_HOOK_OUTPUT_BYTES, run_command_handler, @@ -29,6 +28,7 @@ class HookEngine: snapshot: HooksSnapshot default_timeout: float | None = None max_output_bytes: int = MAX_HOOK_OUTPUT_BYTES + adapter: HookEnvelopeAdapter = field(default_factory=HookEnvelopeAdapter) async def run( self, @@ -53,7 +53,7 @@ async def run( """ match = self.snapshot.match(invocation) try: - payload = serialize_hook_input( + payload = self.adapter.serialize_input( invocation, transcript_path=transcript_path, agent_transcript_path=agent_transcript_path, @@ -64,7 +64,7 @@ async def run( severity="warning", message=f"Could not project hook invocation: {exc}", ) - return reduce_hook_results( + return self.adapter.to_domain_decision( invocation, (), diagnostics=( @@ -92,7 +92,7 @@ async def run( for handler in match.handlers ) ) - return reduce_hook_results( + return self.adapter.to_domain_decision( invocation, results, diagnostics=( diff --git a/libs/code/deepagents_code/hooks/envelope.py b/libs/code/deepagents_code/hooks/envelope.py new file mode 100644 index 0000000000..7c02d723ac --- /dev/null +++ b/libs/code/deepagents_code/hooks/envelope.py @@ -0,0 +1,78 @@ +"""Canonical boundary between hook domain and wire models.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +from deepagents_code.hooks.projection import project_hook_input, serialize_hook_input +from deepagents_code.hooks.reducer import reduce_hook_results + +if TYPE_CHECKING: + from collections.abc import Iterable + from pathlib import Path + + from deepagents_code.hooks.models.domain import ( + HookDecision, + HookDiagnostic, + HookInvocation, + ) + from deepagents_code.hooks.models.wire import HookWireInput + from deepagents_code.hooks.runner import HandlerResult + + +class HookEnvelopeAdapter: + """Adapt hook invocations and handler output across the wire boundary.""" + + @staticmethod + def to_wire_input( + invocation: HookInvocation, + *, + transcript_path: Path, + agent_transcript_path: Path | None = None, + ) -> HookWireInput: + """Project a domain invocation into its validated wire model. + + Returns: + Validated event-specific wire input. + """ + return project_hook_input( + invocation, + transcript_path=transcript_path, + agent_transcript_path=agent_transcript_path, + ) + + @staticmethod + def serialize_input( + invocation: HookInvocation, + *, + transcript_path: Path, + agent_transcript_path: Path | None = None, + ) -> bytes: + """Serialize a domain invocation for command-handler stdin. + + Returns: + Compact validated JSON bytes. + """ + return serialize_hook_input( + invocation, + transcript_path=transcript_path, + agent_transcript_path=agent_transcript_path, + ) + + @staticmethod + def to_domain_decision( + invocation: HookInvocation, + results: Iterable[HandlerResult], + *, + diagnostics: Iterable[HookDiagnostic] = (), + ) -> HookDecision: + """Normalize ordered wire handler results into a domain decision. + + Returns: + Event-specific normalized domain decision. + """ + return reduce_hook_results( + invocation, + results, + diagnostics=diagnostics, + ) diff --git a/libs/code/deepagents_code/hooks/models/domain.py b/libs/code/deepagents_code/hooks/models/domain.py index 3da136a1bf..7a786d11bf 100644 --- a/libs/code/deepagents_code/hooks/models/domain.py +++ b/libs/code/deepagents_code/hooks/models/domain.py @@ -201,7 +201,7 @@ class PostToolUseEvent(_DomainModel): event: Literal[HookEvent.POST_TOOL_USE] call: ToolCallData - result: ToolMessage | Command[str] + result: Command[str] | ToolMessage duration_ms: int | None = None diff --git a/libs/code/deepagents_code/hooks/runtime.py b/libs/code/deepagents_code/hooks/runtime.py index 351c7a003c..91a8150824 100644 --- a/libs/code/deepagents_code/hooks/runtime.py +++ b/libs/code/deepagents_code/hooks/runtime.py @@ -8,6 +8,7 @@ ) from typing import TYPE_CHECKING +from deepagents_code.hooks.client import HookFulfillmentLedger from deepagents_code.hooks.engine import HookEngine from deepagents_code.hooks.loading import load_hooks_config from deepagents_code.hooks.models.domain import ( @@ -50,6 +51,7 @@ class HooksRuntime: transcripts: TranscriptStore engine: HookEngine cwd: Path + fulfillments: HookFulfillmentLedger @classmethod def create( @@ -86,7 +88,13 @@ def create( user_config_dir = config_dir or DEFAULT_CONFIG_DIR store = TranscriptStore(transcript_root or user_config_dir / "transcripts") engine = HookEngine(snapshot) - return cls(snapshot=snapshot, transcripts=store, engine=engine, cwd=cwd) + return cls( + snapshot=snapshot, + transcripts=store, + engine=engine, + cwd=cwd, + fulfillments=HookFulfillmentLedger(), + ) @property def snapshot_id(self) -> str: diff --git a/libs/code/deepagents_code/hooks/server_middleware.py b/libs/code/deepagents_code/hooks/server_middleware.py index 32d07996e4..6c639ad5d4 100644 --- a/libs/code/deepagents_code/hooks/server_middleware.py +++ b/libs/code/deepagents_code/hooks/server_middleware.py @@ -7,12 +7,14 @@ from __future__ import annotations +import hashlib +import json import time from collections.abc import Mapping, Sequence -from dataclasses import dataclass, field +from dataclasses import dataclass, field, replace from datetime import UTC, datetime, timedelta -from typing import TYPE_CHECKING, Any, NotRequired, TypeVar, cast -from uuid import UUID, uuid4 +from typing import TYPE_CHECKING, Any, Literal, NotRequired, TypeAlias, TypeVar, cast +from uuid import UUID, uuid5 from langchain.agents.middleware.human_in_the_loop import ( ActionRequest, @@ -64,19 +66,31 @@ from langchain.tools.tool_node import ToolCallRequest from langchain_core.messages.tool import ToolCall + from langchain_core.tools import BaseTool from langgraph.runtime import Runtime from deepagents_code.json_types import JsonObject _DEFAULT_DEADLINE = timedelta(seconds=600) _STOP_STATE_KEY = "_hooks_stop_continuation_count" +_PRE_TOOL_STATE_KEY = "_hooks_pre_tool_outcomes" _TASK_TOOL_NAME = "task" +_INVOCATION_NAMESPACE = UUID("f2896d18-cf2a-4e7d-b11a-d5b10fc0e335") + +PreToolBehavior: TypeAlias = Literal["allow", "deny", "none"] + + +class _PreToolState(TypedDict): + behavior: PreToolBehavior + reason: str | None + context: list[str] class ServerHooksState(AgentState[Any]): """Agent state extensions for server-owned hook middleware.""" _hooks_stop_continuation_count: NotRequired[int] + _hooks_pre_tool_outcomes: NotRequired[dict[str, _PreToolState]] class _SessionHookGate(TypedDict): @@ -103,6 +117,7 @@ def __init__( cwd: Path, default_deadline: timedelta = _DEFAULT_DEADLINE, emit_stop: bool = True, + mcp_tools: Sequence[BaseTool] = (), ) -> None: """Initialize middleware. @@ -112,11 +127,44 @@ def __init__( emit_stop: Whether to emit the main-agent `Stop` event from `after_agent`. Subagent graphs set this to `False` so they still wrap tools without firing parent `Stop` handlers. + mcp_tools: MCP tools whose server metadata is needed before tool + execution for compatible hook projection. """ super().__init__() self._cwd = cwd self._default_deadline = default_deadline self._emit_stop = emit_stop + self._mcp_servers = { + name: server + for tool in mcp_tools + if (name := getattr(tool, "name", None)) + and isinstance(name, str) + and (server := _mcp_server_from_tool(tool)) is not None + } + + def after_model( + self, + state: ServerHooksState, + runtime: Runtime[ContextT], + ) -> dict[str, Any]: + """Run `PreToolUse` before downstream HITL middleware. + + Returns: + State update carrying per-tool hook outcomes. + """ + return self._after_model(state, runtime) + + async def aafter_model( + self, + state: ServerHooksState, + runtime: Runtime[ContextT], + ) -> dict[str, Any]: + """Run the async graph path through the same interrupt sequence. + + Returns: + State update carrying per-tool hook outcomes. + """ + return self._after_model(state, runtime) def wrap_tool_call( self, @@ -130,16 +178,16 @@ def wrap_tool_call( """ gate = _session_gate(request.runtime.context) call = _tool_call_data(request) + pre = _pre_tool_outcome(request.state, call) context = _hook_context( request.runtime.context, request.runtime.config, self._cwd ) + if pre.blocked is not None: + return _append_message_text(pre.blocked, pre.context) started_or_blocked = self._maybe_subagent_start(request, call, context, gate) if isinstance(started_or_blocked, ToolMessage): return started_or_blocked request = started_or_blocked - pre = self._maybe_pre_tool_use(call, context, gate, request.runtime.config) - if pre.blocked is not None: - return _append_message_text(pre.blocked, pre.context) started = time.perf_counter() result = handler(request) duration_ms = int((time.perf_counter() - started) * 1000) @@ -163,16 +211,16 @@ async def awrap_tool_call( """ gate = _session_gate(request.runtime.context) call = _tool_call_data(request) + pre = _pre_tool_outcome(request.state, call) context = _hook_context( request.runtime.context, request.runtime.config, self._cwd ) + if pre.blocked is not None: + return _append_message_text(pre.blocked, pre.context) started_or_blocked = self._maybe_subagent_start(request, call, context, gate) if isinstance(started_or_blocked, ToolMessage): return started_or_blocked request = started_or_blocked - pre = self._maybe_pre_tool_use(call, context, gate, request.runtime.config) - if pre.blocked is not None: - return _append_message_text(pre.blocked, pre.context) started = time.perf_counter() result = await handler(request) duration_ms = int((time.perf_counter() - started) * 1000) @@ -240,45 +288,62 @@ def _maybe_subagent_start( ) return _inject_subagent_start_context(request, decision) - def _maybe_pre_tool_use( + def _after_model( self, - call: ToolCallData, - context: HookContext, - gate: _SessionHookGate | None, - config: Mapping[str, Any] | None, - ) -> _PreToolOutcome: + state: ServerHooksState, + runtime: Runtime[ContextT], + ) -> dict[str, Any]: + gate = _session_gate(runtime.context) if not _event_enabled(gate, HookEvent.PRE_TOOL_USE): - return _PreToolOutcome() - decision = _invoke_hook( - context, - PreToolUseEvent(event=HookEvent.PRE_TOOL_USE, call=call), - gate=gate, - config=config, - deadline=self._default_deadline, - ) - decision = _require_decision(decision, PreToolUseDecision) - context_parts = tuple(decision.context) - if not decision.continue_processing: - return _PreToolOutcome( - blocked=_denied_tool_message( - call, - PermissionEffect( - behavior="deny", - reason=decision.stop_reason or "Stopped by PreToolUse hook", - ), - ), - context=context_parts, + return {_PRE_TOOL_STATE_KEY: {}} + message = _last_ai_message(state.get("messages", ())) + if message is None: + return {_PRE_TOOL_STATE_KEY: {}} + context = _hook_context(runtime.context, None, self._cwd) + outcomes: dict[str, _PreToolState] = {} + for tool_call in message.tool_calls: + call = _tool_call_data_from_call( + tool_call, + mcp_server=self._mcp_servers.get(str(tool_call.get("name") or "")), ) - behavior = decision.permission.behavior - if behavior == "deny": - return _PreToolOutcome( - blocked=_denied_tool_message(call, decision.permission), - context=context_parts, + decision = _invoke_hook( + context, + PreToolUseEvent(event=HookEvent.PRE_TOOL_USE, call=call), + gate=gate, + config=None, + deadline=self._default_deadline, ) - if behavior == "ask": - blocked = _ask_permission_via_hitl(call, decision.permission) - return _PreToolOutcome(blocked=blocked, context=context_parts) - return _PreToolOutcome(context=context_parts) + decision = _require_decision(decision, PreToolUseDecision) + behavior: PreToolBehavior = "none" + reason: str | None = None + permission = decision.permission + if not decision.continue_processing or permission.behavior == "deny": + behavior = "deny" + reason = ( + permission.reason + or decision.stop_reason + or "Blocked by PreToolUse hook" + ) + elif permission.behavior == "ask": + blocked = _ask_permission_via_hitl(call, permission) + if blocked is None: + behavior = "allow" + else: + behavior = "deny" + blocked_content = blocked.content + reason = ( + blocked_content + if isinstance(blocked_content, str) + else str(blocked_content) + ) + elif permission.behavior == "allow": + behavior = "allow" + outcomes[call.id] = { + "behavior": behavior, + "reason": reason, + "context": list(decision.context), + } + return {_PRE_TOOL_STATE_KEY: outcomes} def _maybe_post_tool_use( self, @@ -291,7 +356,7 @@ def _maybe_post_tool_use( ) -> ToolMessage | Command[Any]: if not _event_enabled(gate, HookEvent.POST_TOOL_USE): return result - if not isinstance(result, ToolMessage): + if _tool_result_failed(result): return result decision = _invoke_hook( context, @@ -406,6 +471,56 @@ def _event_enabled(gate: _SessionHookGate | None, event: HookEvent) -> bool: return gate is not None and event.value in gate["events"] +def pre_tool_behavior(state: object, tool_call_id: str) -> PreToolBehavior | None: + """Return the replayed PreToolUse permission behavior for one call.""" + outcome = _pre_tool_state(state, tool_call_id) + if outcome is None: + return None + behavior = outcome.get("behavior") + if behavior == "allow": + return "allow" + if behavior == "deny": + return "deny" + if behavior == "none": + return "none" + return None + + +def _pre_tool_state(state: object, tool_call_id: str) -> Mapping[str, object] | None: + if not isinstance(state, Mapping): + return None + raw = state.get(_PRE_TOOL_STATE_KEY) + if not isinstance(raw, Mapping): + return None + outcome = raw.get(tool_call_id) + if not isinstance(outcome, Mapping): + return None + return {str(key): value for key, value in outcome.items()} + + +def _pre_tool_outcome(state: object, call: ToolCallData) -> _PreToolOutcome: + outcome = _pre_tool_state(state, call.id) + if outcome is None: + return _PreToolOutcome() + raw_context = outcome.get("context") + context = ( + tuple(item for item in raw_context if isinstance(item, str)) + if isinstance(raw_context, Sequence) and not isinstance(raw_context, str) + else () + ) + if outcome.get("behavior") != "deny": + return _PreToolOutcome(context=context) + raw_reason = outcome.get("reason") + reason = raw_reason if isinstance(raw_reason, str) else None + return _PreToolOutcome( + blocked=_denied_tool_message( + call, + PermissionEffect(behavior="deny", reason=reason), + ), + context=context, + ) + + def _invoke_hook( context: HookContext, event: ( @@ -423,11 +538,18 @@ def _invoke_hook( if gate is None: msg = "hooks_snapshot_id is required to emit server-owned hook events" raise RuntimeError(msg) + run_id = _run_id(config, context.thread_id) + invocation_id = _invocation_id( + run_id=run_id, + snapshot_id=gate["snapshot_id"], + context=context, + event=event, + ) request = HookInvocationRequest( protocol_version=1, - invocation_id=uuid4(), + invocation_id=invocation_id, snapshot_id=gate["snapshot_id"], - run_id=_run_id(config), + run_id=run_id, invocation=HookInvocation(context=context, event=event), deadline=datetime.now(UTC) + deadline, ) @@ -489,15 +611,64 @@ def _context_mapping(runtime_context: object) -> dict[str, Any]: return result -def _run_id(config: Mapping[str, Any] | None) -> str: +def _run_id(config: Mapping[str, Any] | None, thread_id: str) -> str: if isinstance(config, Mapping): configurable = config.get("configurable") if isinstance(configurable, Mapping): for key in ("run_id", "thread_id"): value = configurable.get(key) + if isinstance(value, UUID): + return str(value) if isinstance(value, str) and value: return value - return str(uuid4()) + return thread_id + + +def _invocation_id( + *, + run_id: str, + snapshot_id: str, + context: HookContext, + event: ( + PreToolUseEvent + | PostToolUseEvent + | StopEvent + | SubagentStartEvent + | SubagentStopEvent + ), +) -> UUID: + identity = { + "run_id": run_id, + "thread_id": context.thread_id, + "snapshot_id": snapshot_id, + "event": event.event.value, + "logical_event": _logical_event_identity(context, event), + } + return uuid5( + _INVOCATION_NAMESPACE, + json.dumps(identity, sort_keys=True, separators=(",", ":")), + ) + + +def _logical_event_identity( + context: HookContext, + event: ( + PreToolUseEvent + | PostToolUseEvent + | StopEvent + | SubagentStartEvent + | SubagentStopEvent + ), +) -> str: + if isinstance(event, PreToolUseEvent | PostToolUseEvent): + return event.call.id + if isinstance(event, SubagentStartEvent): + return event.agent.id + prompt_id = str(context.prompt_id) if context.prompt_id is not None else "" + if isinstance(event, SubagentStopEvent): + return f"{event.agent.id}:{event.continuation_count}:{prompt_id}" + message_hash = hashlib.sha256(event.last_assistant_message.encode()).hexdigest() + return f"{event.continuation_count}:{prompt_id}:{message_hash}" def _config_thread_id(config: Mapping[str, Any] | None) -> str | None: @@ -511,7 +682,17 @@ def _config_thread_id(config: Mapping[str, Any] | None) -> str | None: def _tool_call_data(request: ToolCallRequest) -> ToolCallData: - tool_call = request.tool_call + return _tool_call_data_from_call( + request.tool_call, + mcp_server=_mcp_server_from_tool(request.tool), + ) + + +def _tool_call_data_from_call( + tool_call: Mapping[str, object], + *, + mcp_server: str | None, +) -> ToolCallData: raw_args = tool_call.get("args") args: dict[str, Any] if isinstance(raw_args, dict): @@ -522,7 +703,7 @@ def _tool_call_data(request: ToolCallRequest) -> ToolCallData: id=str(tool_call.get("id") or ""), name=str(tool_call.get("name") or ""), args=cast("JsonObject", args), - mcp_server=_mcp_server_from_tool(request.tool), + mcp_server=mcp_server, ) @@ -616,15 +797,15 @@ def _append_message_text( result: ToolMessage | Command[Any], parts: Sequence[str], ) -> ToolMessage | Command[Any]: - if not parts or not isinstance(result, ToolMessage): + if not parts: return result - return _merge_tool_message_content(result, "\n".join(parts)) + return _append_tool_result_text(result, "\n".join(parts)) def _apply_post_tool_use( - result: ToolMessage, + result: ToolMessage | Command[Any], decision: PostToolUseDecision, -) -> ToolMessage: +) -> ToolMessage | Command[Any]: extras: list[str] = [] if decision.feedback: extras.append("\n".join(decision.feedback)) @@ -634,8 +815,9 @@ def _apply_post_tool_use( extras.append(decision.stop_reason) if not extras: return result - return _merge_tool_message_content( - result, "\n\n".join(part for part in extras if part) + return _append_tool_result_text( + result, + "\n\n".join(part for part in extras if part), ) @@ -643,9 +825,49 @@ def _apply_subagent_stop( result: ToolMessage | Command[Any], decision: SubagentStopDecision, ) -> ToolMessage | Command[Any]: - if not decision.context or not isinstance(result, ToolMessage): + if not decision.context: return result - return _merge_tool_message_content(result, "\n".join(decision.context)) + return _append_tool_result_text(result, "\n".join(decision.context)) + + +def _append_tool_result_text( + result: ToolMessage | Command[Any], + suffix: str, +) -> ToolMessage | Command[Any]: + if isinstance(result, ToolMessage): + return _merge_tool_message_content(result, suffix) + update = result.update + if not isinstance(update, Mapping): + return result + raw_messages = update.get("messages") + if not isinstance(raw_messages, Sequence) or isinstance(raw_messages, str): + return result + changed = False + messages: list[object] = [] + for message in raw_messages: + if isinstance(message, ToolMessage): + messages.append(_merge_tool_message_content(message, suffix)) + changed = True + else: + messages.append(message) + if not changed: + return result + return replace(result, update={**update, "messages": messages}) + + +def _tool_result_failed(result: ToolMessage | Command[Any]) -> bool: + if isinstance(result, ToolMessage): + return result.status == "error" + update = result.update + if not isinstance(update, Mapping): + return False + messages = update.get("messages") + if not isinstance(messages, Sequence) or isinstance(messages, str): + return False + return any( + isinstance(message, ToolMessage) and message.status == "error" + for message in messages + ) def _merge_tool_message_content(result: ToolMessage, suffix: str) -> ToolMessage: @@ -705,14 +927,27 @@ def _tool_result_text(result: ToolMessage | Command[Any]) -> str: if isinstance(result, ToolMessage): content = result.content return content if isinstance(content, str) else str(content) - return "" + update = result.update + if not isinstance(update, Mapping): + return "" + messages = update.get("messages") + if not isinstance(messages, Sequence) or isinstance(messages, str): + return "" + return "\n".join( + str(message.content) for message in messages if isinstance(message, ToolMessage) + ) + + +def _last_ai_message(messages: Sequence[Any]) -> AIMessage | None: + return next( + (message for message in reversed(messages) if isinstance(message, AIMessage)), + None, + ) def _last_assistant_text(messages: Sequence[Any]) -> str: - for message in reversed(messages): - if isinstance(message, AIMessage): - content = message.content - if isinstance(content, str): - return content - return str(content) - return "" + message = _last_ai_message(messages) + if message is None: + return "" + content = message.content + return content if isinstance(content, str) else str(content) diff --git a/libs/code/tests/unit_tests/hooks/test_engine.py b/libs/code/tests/unit_tests/hooks/test_engine.py index 9eaa074919..156bc0694b 100644 --- a/libs/code/tests/unit_tests/hooks/test_engine.py +++ b/libs/code/tests/unit_tests/hooks/test_engine.py @@ -15,6 +15,7 @@ from deepagents_code.approval_mode import ApprovalMode from deepagents_code.hooks import dispatch_hook from deepagents_code.hooks.engine import HookEngine +from deepagents_code.hooks.envelope import HookEnvelopeAdapter from deepagents_code.hooks.migration import migrate_legacy_hooks from deepagents_code.hooks.models.adapters import HOOK_WIRE_INPUT_ADAPTER from deepagents_code.hooks.models.config import HooksConfig @@ -83,6 +84,42 @@ def _agent_transcript_path(tmp_path: Path) -> Path: return tmp_path / "agent.jsonl" +def test_envelope_adapter_matches_projection_and_reduction(tmp_path: Path) -> None: + invocation = _invocation( + tmp_path, + UserPromptSubmitEvent( + event=HookEvent.USER_PROMPT_SUBMIT, + prompt="Keep compatibility", + ), + ) + transcript_path = _transcript_path(tmp_path) + adapter = HookEnvelopeAdapter() + result = HandlerResult( + handler_id="handler-1", + output=HookWireOutput.model_validate( + { + "hookSpecificOutput": { + "hookEventName": "UserPromptSubmit", + "additionalContext": "legacy context", + } + } + ), + ) + + assert adapter.to_wire_input( + invocation, + transcript_path=transcript_path, + ) == project_hook_input(invocation, transcript_path=transcript_path) + assert adapter.serialize_input( + invocation, + transcript_path=transcript_path, + ) == serialize_hook_input(invocation, transcript_path=transcript_path) + assert adapter.to_domain_decision(invocation, [result]) == reduce_hook_results( + invocation, + [result], + ) + + def _invocation(tmp_path: Path, event: HookDomainEvent) -> HookInvocation: agent = getattr(event, "agent", None) if not isinstance(agent, AgentIdentity): @@ -1527,19 +1564,22 @@ def test_permission_request_diagnoses_all_deferred_fields(tmp_path: Path) -> Non } -async def test_engine_runs_handlers_concurrently(tmp_path: Path) -> None: +async def test_engine_reduces_in_config_order_when_completion_is_reversed( + tmp_path: Path, +) -> None: first = tmp_path / "first.txt" second = tmp_path / "second.txt" first_cmd = ( "import json,pathlib,time; " - "time.sleep(0.05); " + "time.sleep(0.08); " f"pathlib.Path({str(first)!r}).write_text('first'); " "print(json.dumps({'continue': False, 'stopReason': 'stop'}))" ) second_cmd = ( - "import pathlib,time; " - "time.sleep(0.05); " - f"pathlib.Path({str(second)!r}).write_text('second')" + "import json,pathlib,time; " + "time.sleep(0.01); " + f"pathlib.Path({str(second)!r}).write_text('second'); " + "print(json.dumps({'continue': False, 'stopReason': 'later'}))" ) snapshot = HooksSnapshot.from_config( _config( diff --git a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py index bc2f5da6ea..f687cc5e63 100644 --- a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py +++ b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py @@ -2,15 +2,24 @@ from __future__ import annotations -from datetime import UTC, datetime +import asyncio +import json +import sys +from datetime import UTC, datetime, timedelta from pathlib import Path from typing import TYPE_CHECKING from unittest.mock import MagicMock from uuid import uuid4 import pytest -from langchain_core.messages import ToolMessage - +from langchain_core.language_models.fake_chat_models import GenericFakeChatModel +from langchain_core.messages import AIMessage, ToolMessage +from langgraph.checkpoint.memory import InMemorySaver +from langgraph.graph import START, StateGraph +from langgraph.types import Command +from pydantic import BaseModel + +from deepagents_code.agent import _should_interrupt_tool_call, create_cli_agent from deepagents_code.approval_mode import ApprovalMode from deepagents_code.hooks.client import fulfill_hook_invocation from deepagents_code.hooks.context import apply_hooks_context @@ -49,15 +58,22 @@ _apply_subagent_stop, _ask_permission_via_hitl, _denied_tool_message, + _invoke_hook, _merge_tool_message_content, _session_gate, ) from deepagents_code.hooks.snapshot import HooksSnapshot if TYPE_CHECKING: + from langchain_core.runnables import RunnableConfig + from deepagents_code._cli_context import CLIContext +class _ReplayState(BaseModel): + completed: bool + + def _request(event: PreToolUseEvent | None = None) -> HookInvocationRequest: invocation = HookInvocation( context=HookContext( @@ -118,6 +134,61 @@ def test_hook_resume_value_validates_identity() -> None: ) +def test_real_checkpointer_resume_replays_stable_hook_identity() -> None: + context = HookContext( + thread_id="thread-1", + cwd=Path("/tmp"), + approval_mode=ApprovalMode.MANUAL, + ) + event = PreToolUseEvent( + event=HookEvent.PRE_TOOL_USE, + call=ToolCallData(id="call-1", name="execute", args={"command": "ls"}), + ) + gate = _session_gate( + { + "hooks_snapshot_id": "snapshot-1", + "hooks_server_events": [HookEvent.PRE_TOOL_USE.value], + } + ) + assert gate is not None + + def invoke_hook(state: _ReplayState) -> dict[str, bool]: + assert state.completed is False + decision = _invoke_hook( + context, + event, + gate=gate, + config={"configurable": {"thread_id": "thread-1"}}, + deadline=timedelta(minutes=1), + ) + assert isinstance(decision, PreToolUseDecision) + return {"completed": decision.permission.behavior == "allow"} + + builder = StateGraph(_ReplayState) + builder.add_node("hook", invoke_hook) + builder.add_edge(START, "hook") + graph = builder.compile(checkpointer=InMemorySaver()) + config: RunnableConfig = {"configurable": {"thread_id": "thread-1"}} + + interrupted = graph.invoke(_ReplayState(completed=False), config) + pending = interrupted["__interrupt__"][0] + request = parse_hook_interrupt_payload(pending.value) + assert request is not None + response = HookInvocationResponse( + protocol_version=1, + invocation_id=request.invocation_id, + snapshot_id=request.snapshot_id, + decision=PreToolUseDecision( + event=HookEvent.PRE_TOOL_USE, + permission=PermissionEffect(behavior="allow"), + ), + ) + + resumed = graph.invoke(Command(resume=build_hook_resume_value(response)), config) + + assert resumed["completed"] is True + + def test_apply_hooks_context_sets_server_events(tmp_path: Path) -> None: config_dir = tmp_path / "config" config_dir.mkdir() @@ -204,6 +275,86 @@ def test_apply_post_tool_use_appends_feedback_and_context() -> None: assert "note" in str(updated.content) +def test_post_tool_use_updates_successful_command_result( + monkeypatch: pytest.MonkeyPatch, +) -> None: + middleware = ServerHooksMiddleware(cwd=Path("/tmp")) + result = Command( + update={ + "messages": [ + ToolMessage( + content="ok", + name="execute", + tool_call_id="c1", + ) + ] + } + ) + invoke = MagicMock( + return_value=PostToolUseDecision( + event=HookEvent.POST_TOOL_USE, + context=["post context"], + ) + ) + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware._invoke_hook", + invoke, + ) + + updated = middleware._maybe_post_tool_use( + ToolCallData(id="c1", name="execute", args={}), + HookContext( + thread_id="thread-1", + cwd=Path("/tmp"), + approval_mode=ApprovalMode.MANUAL, + ), + {"snapshot_id": "snap", "events": frozenset({"PostToolUse"})}, + {"configurable": {"thread_id": "thread-1"}}, + result, + 5, + ) + + assert isinstance(updated, Command) + assert isinstance(updated.update, dict) + message = updated.update["messages"][0] + assert isinstance(message, ToolMessage) + assert "post context" in str(message.content) + invoke.assert_called_once() + + +def test_post_tool_use_skips_failed_tool_message( + monkeypatch: pytest.MonkeyPatch, +) -> None: + middleware = ServerHooksMiddleware(cwd=Path("/tmp")) + result = ToolMessage( + content="failed", + name="execute", + tool_call_id="c1", + status="error", + ) + invoke = MagicMock() + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware._invoke_hook", + invoke, + ) + + updated = middleware._maybe_post_tool_use( + ToolCallData(id="c1", name="execute", args={}), + HookContext( + thread_id="thread-1", + cwd=Path("/tmp"), + approval_mode=ApprovalMode.MANUAL, + ), + {"snapshot_id": "snap", "events": frozenset({"PostToolUse"})}, + {"configurable": {"thread_id": "thread-1"}}, + result, + 5, + ) + + assert updated is result + invoke.assert_not_called() + + def test_append_pretool_context_to_result() -> None: result = ToolMessage(content="ran", tool_call_id="c1", name="execute") updated = _append_message_text(result, ("pre context",)) @@ -212,6 +363,170 @@ def test_append_pretool_context_to_result() -> None: assert "pre context" in str(updated.content) +def _pre_tool_runtime() -> MagicMock: + runtime = MagicMock() + runtime.context = { + "hooks_snapshot_id": "snap", + "hooks_server_events": ["PreToolUse"], + "thread_id": "thread-1", + "approval_mode": "manual", + } + return runtime + + +def _pre_tool_state() -> ServerHooksState: + return { + "messages": [ + AIMessage( + content="", + tool_calls=[ + { + "name": "execute", + "args": {"command": "ls"}, + "id": "call-1", + "type": "tool_call", + } + ], + ) + ] + } + + +def _tool_request(state: ServerHooksState, runtime: MagicMock) -> MagicMock: + request = MagicMock() + request.state = state + request.runtime = runtime + request.tool = None + request.tool_call = { + "name": "execute", + "args": {"command": "ls"}, + "id": "call-1", + "type": "tool_call", + } + return request + + +def test_pre_tool_allow_bypasses_hitl_and_preserves_context( + monkeypatch: pytest.MonkeyPatch, +) -> None: + middleware = ServerHooksMiddleware(cwd=Path("/tmp")) + runtime = _pre_tool_runtime() + state = _pre_tool_state() + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware._invoke_hook", + lambda *_args, **_kwargs: PreToolUseDecision( + event=HookEvent.PRE_TOOL_USE, + permission=PermissionEffect(behavior="allow"), + context=["hook context"], + ), + ) + + update = middleware._after_model(state, runtime) + state["_hooks_pre_tool_outcomes"] = update["_hooks_pre_tool_outcomes"] + request = _tool_request(state, runtime) + handler = MagicMock( + return_value=ToolMessage( + content="ran", + name="execute", + tool_call_id="call-1", + ) + ) + + assert _should_interrupt_tool_call(request) is False + result = middleware.wrap_tool_call(request, handler) + assert isinstance(result, ToolMessage) + assert "hook context" in str(result.content) + handler.assert_called_once_with(request) + + +def test_server_pre_tool_node_runs_before_stock_hitl(tmp_path: Path) -> None: + model = GenericFakeChatModel(messages=iter([AIMessage(content="done")])) + model.profile = {"max_input_tokens": 200000} + graph, _backend = create_cli_agent( + model, + "hooks-order-test", + cwd=tmp_path, + enable_memory=False, + enable_skills=False, + enable_shell=False, + ) + edges = {(edge.source, edge.target) for edge in graph.get_graph().edges} + + assert ("model", "ServerHooksMiddleware.after_model") in edges + assert ( + "ServerHooksMiddleware.after_model", + "HumanInTheLoopMiddleware.after_model", + ) in edges + + +def test_pre_tool_ask_reaches_hitl_before_execution( + monkeypatch: pytest.MonkeyPatch, +) -> None: + middleware = ServerHooksMiddleware(cwd=Path("/tmp")) + runtime = _pre_tool_runtime() + state = _pre_tool_state() + order: list[str] = [] + + def invoke(*_args: object, **_kwargs: object) -> PreToolUseDecision: + order.append("hook") + return PreToolUseDecision( + event=HookEvent.PRE_TOOL_USE, + permission=PermissionEffect(behavior="ask", reason="review"), + ) + + def ask(*_args: object, **_kwargs: object) -> None: + order.append("hitl") + + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware._invoke_hook", + invoke, + ) + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware._ask_permission_via_hitl", + ask, + ) + + update = middleware._after_model(state, runtime) + state["_hooks_pre_tool_outcomes"] = update["_hooks_pre_tool_outcomes"] + request = _tool_request(state, runtime) + + assert order == ["hook", "hitl"] + assert _should_interrupt_tool_call(request) is False + + +def test_pre_tool_deny_skips_hitl_and_execution( + monkeypatch: pytest.MonkeyPatch, +) -> None: + middleware = ServerHooksMiddleware(cwd=Path("/tmp")) + runtime = _pre_tool_runtime() + state = _pre_tool_state() + ask = MagicMock() + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware._invoke_hook", + lambda *_args, **_kwargs: PreToolUseDecision( + event=HookEvent.PRE_TOOL_USE, + permission=PermissionEffect(behavior="deny", reason="blocked"), + ), + ) + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware._ask_permission_via_hitl", + ask, + ) + + update = middleware._after_model(state, runtime) + state["_hooks_pre_tool_outcomes"] = update["_hooks_pre_tool_outcomes"] + request = _tool_request(state, runtime) + handler = MagicMock() + + assert _should_interrupt_tool_call(request) is False + result = middleware.wrap_tool_call(request, handler) + assert isinstance(result, ToolMessage) + assert result.status == "error" + assert "blocked" in str(result.content) + ask.assert_not_called() + handler.assert_not_called() + + def test_ask_permission_via_hitl_approve(monkeypatch: pytest.MonkeyPatch) -> None: call = ToolCallData(id="c1", name="execute", args={"command": "ls"}) @@ -347,6 +662,57 @@ async def test_fulfill_hook_invocation_runs_engine(tmp_path: Path) -> None: assert response.decision.permission.behavior in {"allow", "none"} +async def test_fulfillment_is_idempotent_in_flight_and_after_completion( + tmp_path: Path, + caplog: pytest.LogCaptureFixture, +) -> None: + config_dir = tmp_path / "config" + config_dir.mkdir() + marker = tmp_path / "marker.txt" + script = ( + "import json,pathlib,time; " + f"pathlib.Path({str(marker)!r}).write_text('x'); " + "time.sleep(0.05); " + "print(json.dumps({'systemMessage':'once'}))" + ) + (config_dir / "hooks.json").write_text( + json.dumps( + { + "hooks": { + "PreToolUse": [ + { + "hooks": [ + { + "type": "command", + "command": ( + f"{sys.executable} -c {json.dumps(script)}" + ), + } + ] + } + ] + } + } + ), + encoding="utf-8", + ) + runtime = HooksRuntime.create(cwd=tmp_path, config_dir=config_dir) + request = _request().model_copy(update={"snapshot_id": runtime.snapshot_id}) + + with caplog.at_level("WARNING", logger="deepagents_code.hooks.client"): + first, second = await asyncio.gather( + fulfill_hook_invocation(runtime, request), + fulfill_hook_invocation(runtime, request), + ) + third = await fulfill_hook_invocation(runtime, request) + + assert first == second == third + assert marker.read_text() == "x" + assert [record.message for record in caplog.records].count( + "Hook user notice: once" + ) == 1 + + def test_snapshot_configured_server_events() -> None: config = HOOKS_CONFIG_ADAPTER.validate_python( { From 1f00788b1756c1a3de4e785f241bf4218e09943d Mon Sep 17 00:00:00 2001 From: Johannes du Plessis Date: Fri, 24 Jul 2026 12:10:21 -0700 Subject: [PATCH 13/14] fix(code): complete Hooks v2 client lifecycle wiring Populate production transcripts and place prompt and compaction compatibility events at replay-safe lifecycle boundaries across interactive and headless clients. Co-authored-by: Cursor --- .gitignore | 3 - libs/code/deepagents_code/app.py | 138 +++++-- .../deepagents_code/client/non_interactive.py | 305 +++++++++++----- .../deepagents_code/hooks/capabilities.py | 2 +- .../deepagents_code/hooks/client_lifecycle.py | 136 ++++++- libs/code/deepagents_code/hooks/migration.py | 5 +- libs/code/deepagents_code/hooks/projection.py | 15 +- .../hooks/server_middleware.py | 154 ++++++-- libs/code/deepagents_code/hooks/snapshot.py | 3 +- libs/code/deepagents_code/hooks/transcript.py | 97 ++++- .../deepagents_code/tui/textual_adapter.py | 197 ++++++---- .../unit_tests/hooks/test_client_lifecycle.py | 85 +++++ .../unit_tests/hooks/test_configuration.py | 7 +- .../tests/unit_tests/hooks/test_engine.py | 2 +- .../unit_tests/hooks/test_server_lifecycle.py | 339 +++++++++++++++++- .../tests/unit_tests/hooks/test_transcript.py | 74 +++- libs/code/tests/unit_tests/test_app.py | 58 ++- .../tests/unit_tests/test_non_interactive.py | 237 +++++++++++- libs/code/tests/unit_tests/test_offload.py | 160 +++++++-- .../unit_tests/tui/test_textual_adapter.py | 254 +++++++++++++ 20 files changed, 1978 insertions(+), 293 deletions(-) diff --git a/.gitignore b/.gitignore index 6f8e48950f..27be2ced1e 100644 --- a/.gitignore +++ b/.gitignore @@ -228,9 +228,6 @@ __marimo__/ # LangGraph .langgraph_api -# Deep Agents local runtime state -.deepagents/ - #claude .claude diff --git a/libs/code/deepagents_code/app.py b/libs/code/deepagents_code/app.py index a40bb16c6a..8c924d232e 100644 --- a/libs/code/deepagents_code/app.py +++ b/libs/code/deepagents_code/app.py @@ -598,7 +598,10 @@ class _ConfigWriteResult: ClientHookContext, ClientHookService, ) - from deepagents_code.hooks.models.domain import SessionEndCause, SessionStartCause + from deepagents_code.hooks.models.domain import ( + SessionEndCause, + SessionStartCause, + ) from deepagents_code.hooks.runtime import HooksRuntime from deepagents_code.mcp_tools import MCPServerInfo from deepagents_code.model_config import MissingProviderPackageError @@ -1830,6 +1833,9 @@ class _ThreadHistoryPayload: model_params: dict[str, Any] | None = None """Persisted `_model_params` from the checkpoint, if any.""" + transcript_messages: tuple[BaseMessage, ...] = () + """Validated checkpoint messages for Hooks transcript materialization.""" + rubric: str | None = None """Legacy persisted rubric or graph rubric input, if any.""" @@ -3485,16 +3491,16 @@ def __init__( """ self._initial_session_started = False - """Set on first entry into `_run_session_start_sequence` past gating. + """Set after the first client `SessionStart` succeeds. Server respawns (`/mcp reconnect`, `/restart`) post a fresh `ServerReady`; without this flag the sequence re-runs and - `_load_thread_history` bulk-mounts widgets whose IDs already exist in - the DOM, raising `DuplicateIds`. Set on entry (not on success) because - if `_load_thread_history` partially mounted before failing, retrying - would still hit the duplicate-ID path. + `_load_thread_history` bulk-mounts duplicate widgets. """ + self._initial_session_start_stopped = False + """Keep a startup hook stop durable across later `ServerReady` events.""" + # Message queue & store self._pending_messages: deque[QueuedMessage] = deque() """User message queue for sequential processing.""" @@ -8197,6 +8203,8 @@ async def _run_session_start_sequence(self) -> None: await self._auto_accept_pending_goal_rubric() await self._drain_startup_backlog() return + if self._initial_session_start_stopped or self._startup_sequence_running: + return if self._launch_init_requested: self._ensure_launch_init_task() @@ -8211,7 +8219,6 @@ async def _run_session_start_sequence(self) -> None: # worker scheduling (the worker never progresses to set the Event). await self._init_session_state() - self._initial_session_started = True self._startup_sequence_running = True initial_submitted = False try: @@ -8222,8 +8229,6 @@ async def _run_session_start_sequence(self) -> None: if self._initial_resume_requested else SessionStartCause.STARTUP ) - if not await self._run_session_start_hook(start_cause): - return should_load_history = bool(self._lc_thread_id and self._agent) and ( self._resume_thread_intent is not None or not self._has_initial_submission() @@ -8249,6 +8254,11 @@ async def _run_session_start_sequence(self) -> None: ) return + if not await self._run_session_start_hook(start_cause): + self._initial_session_start_stopped = True + return + self._initial_session_started = True + if self._startup_cmd: cmd = self._startup_cmd # One-shot: clear to avoid re-running on any subsequent server swap. @@ -9188,10 +9198,6 @@ async def on_chat_input_submitted(self, event: ChatInput.Submitted) -> None: # Reset quit pending state on any input self._quit_pending = False - from deepagents_code.hooks import dispatch_hook - - await dispatch_hook("user.prompt", {}) - # A bare `exit` quits the app (REPL convention), mirroring `/quit`. # Gated to this interactive path only, so external/scripted callers # (on_external_input) can still send the literal "exit" to the agent. @@ -13463,6 +13469,8 @@ async def _handle_offload(self) -> None: """ from langchain_core.messages.utils import count_tokens_approximately + from deepagents_code.hooks.client_lifecycle import ClientHookStopError + if not self._agent or not self._lc_thread_id: await self._mount_message( AppMessage("Nothing to offload \u2014 start a conversation first"), @@ -13492,11 +13500,6 @@ async def _handle_offload(self) -> None: # Prevent concurrent user input while offload modifies state self._set_agent_running(True) try: - from deepagents_code.hooks import dispatch_hook - - await dispatch_hook("context.offload", {}) - # Keep old hook name for backward compatibility - await dispatch_hook("context.compact", {}) await self._set_spinner("Offloading") prior_event = state_values.get("_summarization_event") @@ -13514,6 +13517,8 @@ async def _handle_offload(self) -> None: tool_error = await self._drive_server_side_compaction( config, seed_tool_call_id ) + except ClientHookStopError: + return except Exception as stream_error: # A server graph can checkpoint the tool-node update before a # later stream transport failure reaches this client. Reconcile @@ -13630,10 +13635,6 @@ async def _handle_offload(self) -> None: f"Context: {before} → {after} tokens " f"({pct}% decrease), {messages_kept} messages kept." ) - from deepagents_code.hooks.models.domain import SessionStartCause - - if not await self._run_session_start_hook(SessionStartCause.COMPACT): - return if archive_path: from deepagents_code.offload import offload_storage_is_ephemeral @@ -13722,6 +13723,11 @@ async def _drive_server_side_compaction( from langgraph.types import Command from deepagents_code.config import settings + from deepagents_code.hooks.client import fulfill_hook_interrupt + from deepagents_code.hooks.client_lifecycle import ClientHookStopError + from deepagents_code.hooks.context import apply_hooks_context + from deepagents_code.hooks.interrupt import is_hook_interrupt_payload + from deepagents_code.hooks.models.domain import SessionStartCause from deepagents_code.offload_middleware import ( COMPACTION_FAILURE_PREFIX, _offload_seed_message_id, @@ -13759,6 +13765,21 @@ async def _drive_server_side_compaction( streaming_agent = cast("Any", agent) seeded_compaction_approved = False + compact_boundary_fired = False + stream_context = CLIContext( + model=self._effective_model_spec(), + model_params=self._model_params_override or {}, + profile_overrides=self._profile_override or {}, + model_context_limit=settings.model_context_limit, + thread_id=self._lc_thread_id, + offload_tool_call_id=tool_call_id, + ) + state = self._session_state + apply_hooks_context( + stream_context, + state.hooks_runtime if state is not None else None, + prompt_id=state.turn_id if state is not None else None, + ) def _decisions_for_interrupt(interrupt_obj: Any) -> list[Any]: # noqa: ANN401 """Approve the forced compaction; reject any other gated tool call. @@ -13829,22 +13850,19 @@ async def _drain(stream_input: Any) -> list[tuple[str, dict[str, Any]]]: # noqa Returns: `(interrupt_id, resume_value)` pairs for every interrupt surfaced during this stream. + + Raises: + ClientHookStopError: If compact-session startup is blocked. + RuntimeError: If a hook interrupt cannot be fulfilled. """ - nonlocal tool_error + nonlocal compact_boundary_fired, tool_error pending: list[tuple[str, dict[str, Any]]] = [] async for chunk in streaming_agent.astream( stream_input, stream_mode=["messages", "updates"], subgraphs=True, config=config, - context=CLIContext( - model=self._effective_model_spec(), - model_params=self._model_params_override or {}, - profile_overrides=self._profile_override or {}, - model_context_limit=settings.model_context_limit, - thread_id=self._lc_thread_id, - offload_tool_call_id=tool_call_id, - ), + context=stream_context, durability="exit", ): if not isinstance(chunk, tuple) or len(chunk) != 3: # noqa: PLR2004 # (namespace, mode, data) @@ -13854,14 +13872,44 @@ async def _drain(stream_input: Any) -> list[tuple[str, dict[str, Any]]]: # noqa for interrupt_obj in data.get("__interrupt__") or []: iid = getattr(interrupt_obj, "id", None) if iid: + value = getattr(interrupt_obj, "value", None) + if is_hook_interrupt_payload(value): + if state is None or state.hooks_runtime is None: + msg = ( + "Received hook invocation interrupt without " + "a HooksRuntime" + ) + raise RuntimeError(msg) + resume = await fulfill_hook_interrupt( + state.hooks_runtime, + value, + ) + if resume is None: + msg = "Failed to parse hook interrupt" + raise RuntimeError(msg) + pending.append((iid, resume)) + continue decisions = _decisions_for_interrupt(interrupt_obj) pending.append((iid, {"decisions": decisions})) elif mode == "messages" and isinstance(data, tuple): msg = data[0] if _is_tool_message(msg): text = _message_text(msg) - if text.startswith(COMPACTION_FAILURE_PREFIX): + if text.startswith(COMPACTION_FAILURE_PREFIX) or ( + getattr(msg, "name", None) == "compact_conversation" + and getattr(msg, "status", None) == "error" + ): tool_error = text + elif ( + text.startswith("Conversation compacted.") + and not compact_boundary_fired + ): + compact_boundary_fired = True + if not await self._run_session_start_hook( + SessionStartCause.COMPACT + ): + msg = "Compact continuation stopped by hook" + raise ClientHookStopError(msg) return pending # Bound the resume loop: after compaction the model runs again, and a @@ -15107,7 +15155,16 @@ async def _fetch_thread_history_data(self, thread_id: str) -> _ThreadHistoryPayl # Offload conversion so large histories don't block the UI loop. data = await asyncio.to_thread(self._convert_messages_to_data, messages) - return replace(payload, messages=data) + from langchain_core.messages import BaseMessage + + transcript_messages = tuple( + message for message in messages if isinstance(message, BaseMessage) + ) + return replace( + payload, + messages=data, + transcript_messages=transcript_messages, + ) async def _adopt_resumed_model_if_needed( self, @@ -15257,6 +15314,16 @@ async def _load_thread_history( else await self._fetch_thread_history_data(history_thread_id) ) self._restore_goal_rubric_state(payload) + state = self._session_state + if ( + state is not None + and state.hooks_runtime is not None + and payload.transcript_messages + ): + state.hooks_runtime.append_messages( + history_thread_id, + payload.transcript_messages, + ) # Adopt the resumed thread's model (session-only) so the session # continues on the model it was last using, not the global default. @@ -22421,6 +22488,7 @@ async def _resume_thread(self, thread_id: str) -> None: # choice for this session. Consumed by `_load_thread_history`. self._should_adopt_resumed_model = not self._model_explicitly_set + await self._refresh_client_hooks_runtime() # Load thread history await self._load_thread_history( thread_id=thread_id, @@ -22433,7 +22501,6 @@ async def _resume_thread(self, thread_id: str) -> None: # thread". Set only after the last statement that can raise, so a # failed switch (handled below) never leaves a stale pointer. self._session_state.previous_thread_id = prev_session_thread - await self._refresh_client_hooks_runtime() if not await self._run_session_start_hook(SessionStartCause.RESUME): return except Exception as exc: @@ -22461,6 +22528,8 @@ async def _resume_thread(self, thread_id: str) -> None: ) await self._restore_cwd_after_failed_thread_switch(prev_cwd) rollback_restore_failed = False + if outgoing_ended: + await self._refresh_client_hooks_runtime() # Attempt to restore the previous thread's visible history try: await self._clear_messages() @@ -22473,7 +22542,6 @@ async def _resume_thread(self, thread_id: str) -> None: ) logger.warning(msg, thread_id, exc_info=True) if outgoing_ended: - await self._refresh_client_hooks_runtime() await self._run_session_start_hook(SessionStartCause.RESUME) error_message = f"Failed to switch to thread {thread_id}: {exc}." if rollback_restore_failed: diff --git a/libs/code/deepagents_code/client/non_interactive.py b/libs/code/deepagents_code/client/non_interactive.py index c2b1f23558..2ec178af2d 100644 --- a/libs/code/deepagents_code/client/non_interactive.py +++ b/libs/code/deepagents_code/client/non_interactive.py @@ -29,7 +29,7 @@ from typing import TYPE_CHECKING, Any, NoReturn, cast from langchain.agents.middleware.human_in_the_loop import ActionRequest, HITLRequest -from langchain_core.messages import AIMessage, ToolMessage +from langchain_core.messages import AIMessage, HumanMessage, ToolMessage from langgraph.types import Command, Interrupt from pydantic import TypeAdapter, ValidationError from rich.console import Console @@ -89,8 +89,13 @@ from langchain_core.runnables import RunnableConfig from deepagents_code.approval_mode import ApprovalMode - from deepagents_code.hooks.client_lifecycle import ClientHookService + from deepagents_code.hooks.client_lifecycle import ( + ClientHookContext, + ClientHookService, + ) + from deepagents_code.hooks.models.domain import SessionEndCause from deepagents_code.hooks.runtime import HooksRuntime + from deepagents_code.hooks.transcript import TranscriptRecorder logger = logging.getLogger(__name__) @@ -108,6 +113,12 @@ def _raise_hitl_iteration_limit(message: str) -> NoReturn: raise HITLIterationLimitError(message) +def _raise_client_hook_stop(message: str) -> NoReturn: + from deepagents_code.hooks.client_lifecycle import ClientHookStopError + + raise ClientHookStopError(message) + + _HITL_REQUEST_ADAPTER = TypeAdapter(HITLRequest) _STREAM_CHUNK_LENGTH = 3 @@ -380,9 +391,24 @@ class StreamState: client_hooks: ClientHookService | None = None """Client-owned lifecycle facade for this headless session.""" + client_hook_context: ClientHookContext | None = None + """Validated context shared by headless client lifecycle events.""" + + transcript: TranscriptRecorder | None = None + """Completed root and identified subagent stream messages.""" + + active_model: str | None = None + """Model projected into compact lifecycle events.""" + summarization_observed: bool = False """Whether the current stream crossed a compaction boundary.""" + completed_compaction_ids: set[str] = field(default_factory=set) + """Compaction tool results whose post-boundary lifecycle already fired.""" + + session_end_fired: bool = False + """Whether the headless client session has emitted its terminal event.""" + interrupt_occurred: bool = False """Flag indicating whether any HITL interrupt was received during the current stream pass.""" @@ -852,6 +878,24 @@ def _process_stream_chunk( namespace, stream_mode, data = chunk is_main_agent = not namespace + if ( + stream_mode == "messages" + and isinstance(data, tuple) + and len(data) == _MESSAGE_DATA_LENGTH + and state.transcript is not None + ): + message, metadata = data + transcript_metadata = ( + {str(key): value for key, value in metadata.items()} + if isinstance(metadata, dict) + else None + ) + state.transcript.record( + message, + transcript_metadata, + main_agent=is_main_agent, + ) + if not is_main_agent: return @@ -1015,23 +1059,29 @@ async def _process_hitl_interrupts( current_interrupts = dict(state.pending_interrupts) state.pending_interrupts.clear() - from deepagents_code.approval_mode import ApprovalMode from deepagents_code.hooks.client_lifecycle import ( ClientHookContext, ClientHookStopError, + PermissionReviewDecision, + permission_hook_outcome, + permission_review_payload, ) from deepagents_code.hooks.models.domain import ( DcodeNotificationKind, ToolCallData, ) - context = ClientHookContext.create( - thread_id=thread_id, - approval_mode=ApprovalMode.MANUAL, - ) + context = state.client_hook_context + if context is None: + from deepagents_code.approval_mode import ApprovalMode + + context = ClientHookContext.create( + thread_id=thread_id, + approval_mode=ApprovalMode.MANUAL, + ) for interrupt_id, hitl_request in current_interrupts.items(): action_requests = hitl_request["action_requests"] - decisions: list[dict[str, str] | None] = [] + decisions: list[PermissionReviewDecision | None] = [] for index, action_request in enumerate(action_requests): try: hook_decision = ( @@ -1055,23 +1105,15 @@ async def _process_hitl_interrupts( if hook_decision is None: decisions.append(None) continue - if not hook_decision.continue_processing: - reason = hook_decision.stop_reason or "Permission stopped by hook" - raise ClientHookStopError(reason) - permission = hook_decision.permission - if permission.behavior == "allow": - decisions.append({"type": "approve"}) - elif permission.behavior == "deny": - denied = {"type": "reject"} - if permission.reason: - denied["message"] = permission.reason - if permission.interrupt: - raise ClientHookStopError( - permission.reason or "Permission interrupted by hook" - ) - decisions.append(denied) - else: - decisions.append(None) + outcome = permission_hook_outcome(hook_decision) + if outcome.interrupt: + reason = ( + outcome.decision.get("message") + if outcome.decision is not None + else None + ) + raise ClientHookStopError(reason or "Permission interrupted by hook") + decisions.append(outcome.decision) if ( any(decision is None for decision in decisions) @@ -1094,7 +1136,7 @@ async def _process_hitl_interrupts( strict=True, ): resolved.append( - decision + permission_review_payload(decision) if decision is not None else _make_hitl_decision(action_request, console) ) @@ -1133,12 +1175,102 @@ async def _stream_agent( context=context, durability="exit", ): + summarization = _summarization_stream_status(chunk) + compaction_id = _compaction_result_id(chunk) + if summarization is False and state.summarization_observed: + if compaction_id is not None: + state.completed_compaction_ids.add(compaction_id) + await _after_headless_compact(state) + state.summarization_observed = False + elif ( + compaction_id is not None + and compaction_id not in state.completed_compaction_ids + ): + state.completed_compaction_ids.add(compaction_id) + await _after_headless_compact(state) _process_stream_chunk(chunk, state, console, file_op_tracker) + if state.summarization_observed: + await _after_headless_compact(state) + state.summarization_observed = False finally: if state.spinner: state.spinner.stop() +def _summarization_stream_status(chunk: object) -> bool | None: + if not isinstance(chunk, tuple) or len(chunk) != _STREAM_CHUNK_LENGTH: + return None + namespace, stream_mode, data = chunk + if namespace: + return None + if ( + stream_mode != "messages" + or not isinstance(data, tuple) + or len(data) != _MESSAGE_DATA_LENGTH + ): + return None + _message, metadata = data + return isinstance(metadata, dict) and metadata.get("lc_source") == "summarization" + + +def _compaction_result_id(chunk: object) -> str | None: + if not isinstance(chunk, tuple) or len(chunk) != _STREAM_CHUNK_LENGTH: + return None + namespace, stream_mode, data = chunk + if namespace: + return None + if ( + stream_mode != "messages" + or not isinstance(data, tuple) + or len(data) != _MESSAGE_DATA_LENGTH + ): + return None + message, _metadata = data + if not ( + isinstance(message, ToolMessage) + and str(message.content).startswith("Conversation compacted.") + ): + return None + tool_call_id = getattr(message, "tool_call_id", None) + return tool_call_id if isinstance(tool_call_id, str) and tool_call_id else None + + +async def _after_headless_compact(state: StreamState) -> None: + from deepagents_code.hooks.client_lifecycle import ClientHookStopError + from deepagents_code.hooks.models.domain import SessionStartCause + + service = state.client_hooks + context = state.client_hook_context + if service is None or context is None: + return + try: + decision = await service.session_start( + context, + SessionStartCause.COMPACT, + model=state.active_model, + ) + except Exception: + logger.warning("Compact SessionStart hook invocation failed", exc_info=True) + return + if not decision.continue_processing: + reason = decision.stop_reason or "Compact session start stopped by hook" + raise ClientHookStopError(reason) + + +async def _end_headless_session( + state: StreamState, + context: ClientHookContext, + cause: SessionEndCause, +) -> None: + if state.client_hooks is None or state.session_end_fired: + return + state.session_end_fired = True + try: + await state.client_hooks.session_end(context, cause) + except Exception: + logger.warning("SessionEnd hook invocation failed", exc_info=True) + + def _dispatch_orphaned_tool_result_hooks(state: StreamState, tool_output: str) -> None: """Close out `tool.use` events that never received a `ToolMessage`. @@ -1313,6 +1445,12 @@ async def _run_agent_loop( approval_mode=resolved_approval_mode, prompt_id=prompt_id, ) + state.client_hook_context = client_context + state.active_model = settings.model_name or None + if hooks_runtime is not None: + from deepagents_code.hooks.transcript import TranscriptRecorder + + state.transcript = TranscriptRecorder(hooks_runtime, thread_id) if state.client_hooks is not None: try: start_decision = await state.client_hooks.session_start( @@ -1329,13 +1467,11 @@ async def _run_agent_loop( session_context = state.client_hooks.take_session_context(thread_id) if start_decision is not None and not start_decision.continue_processing: reason = start_decision.stop_reason or "Session start stopped by hook" - try: - await state.client_hooks.session_end( - client_context, - SessionEndCause.OTHER, - ) - except Exception: - logger.warning("SessionEnd hook invocation failed", exc_info=True) + await _end_headless_session( + state, + client_context, + SessionEndCause.OTHER, + ) raise ClientHookStopError(reason) if session_context: messages = stream_input["messages"] @@ -1347,7 +1483,43 @@ async def _run_agent_loop( }, ) - await dispatch_hook("session.start", {"thread_id": thread_id}) + try: + if state.transcript is not None: + state.transcript.append([HumanMessage(content=message)]) + if state.client_hooks is not None and state.client_hooks.has_handlers( + HookEvent.USER_PROMPT_SUBMIT + ): + prompt_decision = await state.client_hooks.user_prompt_submit( + client_context, + message, + ) + if not prompt_decision.continue_processing: + reason = ( + prompt_decision.stop_reason + or "User prompt submission stopped by hook" + ) + _raise_client_hook_stop(reason) + messages = stream_input["messages"] + if prompt_decision.context: + messages.insert( + len(messages) - 1, + { + "role": "system", + "content": "\n\n".join(prompt_decision.context), + }, + ) + if prompt_decision.suppress_original_prompt: + messages.pop() + else: + await dispatch_hook("session.start", {"thread_id": thread_id}) + await dispatch_hook("user.prompt", {}) + except BaseException: + await _end_headless_session( + state, + client_context, + SessionEndCause.OTHER, + ) + raise start_time = time.monotonic() @@ -1356,26 +1528,6 @@ async def _run_agent_loop( await _stream_agent( agent, stream_input, config, state, console, file_op_tracker, context ) - if state.summarization_observed and state.client_hooks is not None: - try: - compact_decision = await state.client_hooks.session_start( - client_context, - SessionStartCause.COMPACT, - model=settings.model_name or None, - ) - except Exception: - logger.warning( - "Compact SessionStart hook invocation failed", - exc_info=True, - ) - else: - if not compact_decision.continue_processing: - reason = ( - compact_decision.stop_reason - or "Compact session start stopped by hook" - ) - raise ClientHookStopError(reason) - state.summarization_observed = False # The internal default applies when --max-turns is omitted, guarding # against unbounded runaway loops in scripts that forgot to set one. @@ -1409,35 +1561,12 @@ async def _run_agent_loop( await _stream_agent( agent, stream_input, config, state, console, file_op_tracker, context ) - if state.summarization_observed and state.client_hooks is not None: - try: - compact_decision = await state.client_hooks.session_start( - client_context, - SessionStartCause.COMPACT, - model=settings.model_name or None, - ) - except Exception: - logger.warning( - "Compact SessionStart hook invocation failed", - exc_info=True, - ) - else: - if not compact_decision.continue_processing: - reason = ( - compact_decision.stop_reason - or "Compact session start stopped by hook" - ) - raise ClientHookStopError(reason) - state.summarization_observed = False except BaseException: - if state.client_hooks is not None: - try: - await state.client_hooks.session_end( - client_context, - SessionEndCause.OTHER, - ) - except Exception: - logger.warning("SessionEnd hook invocation failed", exc_info=True) + await _end_headless_session( + state, + client_context, + SessionEndCause.OTHER, + ) raise finally: # Close out any `tool.use` with no matching `ToolMessage` — e.g. a stream @@ -1521,13 +1650,11 @@ async def _run_agent_loop( logger.warning("Notification hook invocation failed", exc_info=True) if not state.client_hooks.has_handlers(HookEvent.NOTIFICATION): await dispatch_hook("task.complete", {"thread_id": thread_id}) - try: - await state.client_hooks.session_end( - client_context, - SessionEndCause.PROMPT_INPUT_EXIT, - ) - except Exception: - logger.warning("SessionEnd hook invocation failed", exc_info=True) + await _end_headless_session( + state, + client_context, + SessionEndCause.PROMPT_INPUT_EXIT, + ) if not state.client_hooks.has_handlers(HookEvent.SESSION_END): await dispatch_hook("session.end", {"thread_id": thread_id}) if notification_stop is not None: diff --git a/libs/code/deepagents_code/hooks/capabilities.py b/libs/code/deepagents_code/hooks/capabilities.py index 1d4fe3705a..5247efa96d 100644 --- a/libs/code/deepagents_code/hooks/capabilities.py +++ b/libs/code/deepagents_code/hooks/capabilities.py @@ -185,7 +185,7 @@ class HookEventSpec: ), HookEvent.PRE_COMPACT: HookEventSpec( event=HookEvent.PRE_COMPACT, - owner=HookOwner.CLIENT, + owner=HookOwner.SERVER, event_model=PreCompactEvent, decision_model=PreCompactDecision, matcher_field="trigger", diff --git a/libs/code/deepagents_code/hooks/client_lifecycle.py b/libs/code/deepagents_code/hooks/client_lifecycle.py index aa3134f3f4..c8b39f8f25 100644 --- a/libs/code/deepagents_code/hooks/client_lifecycle.py +++ b/libs/code/deepagents_code/hooks/client_lifecycle.py @@ -5,11 +5,12 @@ import logging import sys from dataclasses import dataclass, field -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Literal, NotRequired, TypedDict from uuid import UUID from deepagents_code.approval_mode import ApprovalMode from deepagents_code.hooks.models.domain import ( + CompactTrigger, DcodeNotification, DcodeNotificationKind, HookContext, @@ -23,6 +24,8 @@ PermissionEffect, PermissionRequestDecision, PermissionRequestEvent, + PreCompactDecision, + PreCompactEvent, SessionEndCause, SessionEndDecision, SessionEndEvent, @@ -30,6 +33,8 @@ SessionStartDecision, SessionStartEvent, ToolCallData, + UserPromptSubmitDecision, + UserPromptSubmitEvent, ) if TYPE_CHECKING: @@ -53,6 +58,69 @@ class ClientHookStopError(RuntimeError): """Raised when a client-owned hook stops lifecycle processing.""" +class PermissionReviewDecision(TypedDict): + """Client approval decision compatible with HITL resume payloads.""" + + type: Literal["approve", "reject"] + message: NotRequired[str] + + +@dataclass(frozen=True, slots=True) +class PermissionHookOutcome: + """Normalized result shared by TUI and headless permission handling.""" + + decision: PermissionReviewDecision | None + interrupt: bool = False + + +def permission_hook_outcome( + decision: PermissionRequestDecision, +) -> PermissionHookOutcome: + """Translate a hook permission decision into a client review outcome. + + Args: + decision: Aggregated permission hook decision. + + Returns: + Shared approval, rejection, or unresolved result. + """ + if not decision.continue_processing: + return PermissionHookOutcome( + { + "type": "reject", + "message": decision.stop_reason or "Permission stopped by hook", + }, + interrupt=True, + ) + permission = decision.permission + if permission.behavior == "allow": + return PermissionHookOutcome({"type": "approve"}) + if permission.behavior == "deny": + denied = PermissionReviewDecision(type="reject") + if permission.reason: + denied["message"] = permission.reason + return PermissionHookOutcome(denied, interrupt=permission.interrupt) + return PermissionHookOutcome(None) + + +def permission_review_payload( + decision: PermissionReviewDecision, +) -> dict[str, str]: + """Copy a typed review decision into a mutable resume payload. + + Args: + decision: Structurally validated hook review decision. + + Returns: + Mutable HITL resume payload. + """ + payload = {"type": decision["type"]} + message = decision.get("message") + if message is not None: + payload["message"] = message + return payload + + @dataclass(frozen=True, slots=True) class ClientHookContext: """Client state required to create a domain hook invocation.""" @@ -175,6 +243,72 @@ async def session_end( self._session_context.pop(context.thread_id, None) return decision + async def user_prompt_submit( + self, + context: ClientHookContext, + prompt: str, + ) -> UserPromptSubmitDecision: + """Invoke `UserPromptSubmit` before a user turn. + + Args: + context: Current client turn context. + prompt: Original user prompt. + + Returns: + Aggregated prompt decision. + + Raises: + TypeError: If the runtime returns a mismatched decision type. + """ + if not self.has_handlers(HookEvent.USER_PROMPT_SUBMIT): + return UserPromptSubmitDecision(event=HookEvent.USER_PROMPT_SUBMIT) + decision = await self._invoke( + context, + UserPromptSubmitEvent( + event=HookEvent.USER_PROMPT_SUBMIT, + prompt=prompt, + ), + ) + if not isinstance(decision, UserPromptSubmitDecision): + msg = f"Expected UserPromptSubmitDecision, got {type(decision).__name__}" + raise TypeError(msg) + return decision + + async def pre_compact( + self, + context: ClientHookContext, + trigger: CompactTrigger, + *, + custom_instructions: str = "", + ) -> PreCompactDecision: + """Invoke `PreCompact` through the session hook runtime. + + Args: + context: Current client turn context. + trigger: Manual or automatic compaction source. + custom_instructions: Optional compaction instructions. + + Returns: + Aggregated pre-compaction decision. + + Raises: + TypeError: If the runtime returns a mismatched decision type. + """ + if not self.has_handlers(HookEvent.PRE_COMPACT): + return PreCompactDecision(event=HookEvent.PRE_COMPACT) + decision = await self._invoke( + context, + PreCompactEvent( + event=HookEvent.PRE_COMPACT, + trigger=trigger, + custom_instructions=custom_instructions, + ), + ) + if not isinstance(decision, PreCompactDecision): + msg = f"Expected PreCompactDecision, got {type(decision).__name__}" + raise TypeError(msg) + return decision + async def permission_request( self, context: ClientHookContext, diff --git a/libs/code/deepagents_code/hooks/migration.py b/libs/code/deepagents_code/hooks/migration.py index e5848b546a..36be2d0e22 100644 --- a/libs/code/deepagents_code/hooks/migration.py +++ b/libs/code/deepagents_code/hooks/migration.py @@ -1,8 +1,7 @@ """Legacy dotted-event migration helpers for Hooks v2 configuration. -These utilities are intentionally not activated at legacy dispatch call sites. -Lifecycle wiring belongs to later tickets; this module only converts -semantically equivalent config when an explicit loader asks for it. +Legacy documents are converted by the loader so lifecycle call sites dispatch +only canonical events and do not duplicate old dotted-event hooks. """ from __future__ import annotations diff --git a/libs/code/deepagents_code/hooks/projection.py b/libs/code/deepagents_code/hooks/projection.py index c2bd2fe26e..7cea2c4199 100644 --- a/libs/code/deepagents_code/hooks/projection.py +++ b/libs/code/deepagents_code/hooks/projection.py @@ -184,7 +184,7 @@ def _project_notification( hook_event_name=HookEvent.NOTIFICATION, message=event.notification.message, title=event.notification.title, - notification_type=_notification_type(event.notification.type), + notification_type=to_wire_notification_type(event.notification.type), ) @@ -361,7 +361,18 @@ def _permission_mode(mode: ApprovalMode) -> WirePermissionMode: }[mode] -def _notification_type(value: str) -> WireNotificationType: +def to_wire_notification_type(value: str) -> WireNotificationType: + """Return the compatible notification matcher and wire value. + + Args: + value: Domain or wire notification type. + + Returns: + Canonical wire notification type. + + Raises: + ValueError: If the notification type is unsupported. + """ mappings: dict[str, WireNotificationType] = { DcodeNotificationKind.PERMISSION_REQUIRED: ( WireNotificationType.PERMISSION_PROMPT diff --git a/libs/code/deepagents_code/hooks/server_middleware.py b/libs/code/deepagents_code/hooks/server_middleware.py index 6c639ad5d4..0f79a716af 100644 --- a/libs/code/deepagents_code/hooks/server_middleware.py +++ b/libs/code/deepagents_code/hooks/server_middleware.py @@ -1,8 +1,8 @@ """Server-owned Hooks v2 lifecycle middleware. -Emits `PreToolUse`, `PostToolUse`, `Stop`, `SubagentStart`, and `SubagentStop` -through the LangGraph interrupt channel so the client runtime can execute -matching handlers and return typed decisions. +Emits `PreCompact`, `PreToolUse`, `PostToolUse`, `Stop`, `SubagentStart`, and +`SubagentStop` through the LangGraph interrupt channel so the client runtime can +execute matching handlers and return typed decisions. """ from __future__ import annotations @@ -11,6 +11,7 @@ import json import time from collections.abc import Mapping, Sequence +from contextlib import contextmanager from dataclasses import dataclass, field, replace from datetime import UTC, datetime, timedelta from typing import TYPE_CHECKING, Any, Literal, NotRequired, TypeAlias, TypeVar, cast @@ -40,6 +41,7 @@ from deepagents_code.hooks.models.domain import ( AgentIdentity, BaseHookDecision, + CompactTrigger, HookContext, HookDecision, HookEvent, @@ -47,6 +49,8 @@ PermissionEffect, PostToolUseDecision, PostToolUseEvent, + PreCompactDecision, + PreCompactEvent, PreToolUseDecision, PreToolUseEvent, StopDecision, @@ -61,11 +65,12 @@ from deepagents_code.hooks.tools import to_wire_tool_name if TYPE_CHECKING: - from collections.abc import Awaitable, Callable + from collections.abc import Awaitable, Callable, Iterator from pathlib import Path from langchain.tools.tool_node import ToolCallRequest from langchain_core.messages.tool import ToolCall + from langchain_core.runnables import RunnableConfig from langchain_core.tools import BaseTool from langgraph.runtime import Runtime @@ -75,6 +80,7 @@ _STOP_STATE_KEY = "_hooks_stop_continuation_count" _PRE_TOOL_STATE_KEY = "_hooks_pre_tool_outcomes" _TASK_TOOL_NAME = "task" +_COMPACT_TOOL_NAME = "compact_conversation" _INVOCATION_NAMESPACE = UUID("f2896d18-cf2a-4e7d-b11a-d5b10fc0e335") PreToolBehavior: TypeAlias = Literal["allow", "deny", "none"] @@ -100,12 +106,38 @@ class _SessionHookGate(TypedDict): @dataclass(slots=True) class _PreToolOutcome: - """PreToolUse gate result for the tool-call wrapper.""" + """Pre-execution gate result for the tool-call wrapper.""" blocked: ToolMessage | None = None context: tuple[str, ...] = field(default_factory=tuple) +@contextmanager +def _subagent_transcript_config( + call: ToolCallData, + config: RunnableConfig, +) -> Iterator[None]: + if call.name != _TASK_TOOL_NAME: + yield + return + + from langchain_core.runnables.config import var_child_runnable_config + + from deepagents_code.hooks.transcript import ( + SUBAGENT_TRANSCRIPT_ID_METADATA_KEY, + ) + + metadata = config.get("metadata") + child_metadata = dict(metadata) if isinstance(metadata, Mapping) else {} + child_metadata[SUBAGENT_TRANSCRIPT_ID_METADATA_KEY] = call.id + child_config: RunnableConfig = {**config, "metadata": child_metadata} + token = var_child_runnable_config.set(child_config) + try: + yield + finally: + var_child_runnable_config.reset(token) + + class ServerHooksMiddleware(AgentMiddleware[ServerHooksState, ContextT, ResponseT]): """Emit server-owned lifecycle events over the hook interrupt transport.""" @@ -147,7 +179,7 @@ def after_model( state: ServerHooksState, runtime: Runtime[ContextT], ) -> dict[str, Any]: - """Run `PreToolUse` before downstream HITL middleware. + """Run pre-execution hooks before downstream HITL middleware. Returns: State update carrying per-tool hook outcomes. @@ -189,7 +221,8 @@ def wrap_tool_call( return started_or_blocked request = started_or_blocked started = time.perf_counter() - result = handler(request) + with _subagent_transcript_config(call, request.runtime.config): + result = handler(request) duration_ms = int((time.perf_counter() - started) * 1000) result = _append_message_text(result, pre.context) result = self._maybe_post_tool_use( @@ -222,7 +255,8 @@ async def awrap_tool_call( return started_or_blocked request = started_or_blocked started = time.perf_counter() - result = await handler(request) + with _subagent_transcript_config(call, request.runtime.config): + result = await handler(request) duration_ms = int((time.perf_counter() - started) * 1000) result = _append_message_text(result, pre.context) result = self._maybe_post_tool_use( @@ -294,7 +328,9 @@ def _after_model( runtime: Runtime[ContextT], ) -> dict[str, Any]: gate = _session_gate(runtime.context) - if not _event_enabled(gate, HookEvent.PRE_TOOL_USE): + precompact_enabled = _event_enabled(gate, HookEvent.PRE_COMPACT) + pretool_enabled = _event_enabled(gate, HookEvent.PRE_TOOL_USE) + if not precompact_enabled and not pretool_enabled: return {_PRE_TOOL_STATE_KEY: {}} message = _last_ai_message(state.get("messages", ())) if message is None: @@ -306,42 +342,67 @@ def _after_model( tool_call, mcp_server=self._mcp_servers.get(str(tool_call.get("name") or "")), ) - decision = _invoke_hook( - context, - PreToolUseEvent(event=HookEvent.PRE_TOOL_USE, call=call), - gate=gate, - config=None, - deadline=self._default_deadline, - ) - decision = _require_decision(decision, PreToolUseDecision) behavior: PreToolBehavior = "none" reason: str | None = None - permission = decision.permission - if not decision.continue_processing or permission.behavior == "deny": - behavior = "deny" - reason = ( - permission.reason - or decision.stop_reason - or "Blocked by PreToolUse hook" + hook_context: list[str] = [] + if precompact_enabled and call.name == _COMPACT_TOOL_NAME: + trigger = ( + CompactTrigger.MANUAL + if call.args.get("force") is True + else CompactTrigger.AUTO ) - elif permission.behavior == "ask": - blocked = _ask_permission_via_hitl(call, permission) - if blocked is None: - behavior = "allow" - else: + compact = _invoke_hook( + context, + PreCompactEvent(event=HookEvent.PRE_COMPACT, trigger=trigger), + gate=gate, + config=None, + deadline=self._default_deadline, + logical_event_id=call.id, + ) + compact = _require_decision(compact, PreCompactDecision) + if not compact.continue_processing: + outcomes[call.id] = { + "behavior": "deny", + "reason": compact.stop_reason or "Blocked by PreCompact hook", + "context": hook_context, + } + continue + if pretool_enabled: + decision = _invoke_hook( + context, + PreToolUseEvent(event=HookEvent.PRE_TOOL_USE, call=call), + gate=gate, + config=None, + deadline=self._default_deadline, + ) + decision = _require_decision(decision, PreToolUseDecision) + permission = decision.permission + hook_context.extend(decision.context) + if not decision.continue_processing or permission.behavior == "deny": behavior = "deny" - blocked_content = blocked.content reason = ( - blocked_content - if isinstance(blocked_content, str) - else str(blocked_content) + permission.reason + or decision.stop_reason + or "Blocked by PreToolUse hook" ) - elif permission.behavior == "allow": - behavior = "allow" + elif permission.behavior == "ask": + blocked = _ask_permission_via_hitl(call, permission) + if blocked is None: + behavior = "allow" + else: + behavior = "deny" + blocked_content = blocked.content + reason = ( + blocked_content + if isinstance(blocked_content, str) + else str(blocked_content) + ) + elif permission.behavior == "allow": + behavior = "allow" outcomes[call.id] = { "behavior": behavior, "reason": reason, - "context": list(decision.context), + "context": hook_context, } return {_PRE_TOOL_STATE_KEY: outcomes} @@ -472,7 +533,7 @@ def _event_enabled(gate: _SessionHookGate | None, event: HookEvent) -> bool: def pre_tool_behavior(state: object, tool_call_id: str) -> PreToolBehavior | None: - """Return the replayed PreToolUse permission behavior for one call.""" + """Return the replayed pre-execution hook behavior for one call.""" outcome = _pre_tool_state(state, tool_call_id) if outcome is None: return None @@ -526,6 +587,7 @@ def _invoke_hook( event: ( PreToolUseEvent | PostToolUseEvent + | PreCompactEvent | StopEvent | SubagentStartEvent | SubagentStopEvent @@ -534,6 +596,7 @@ def _invoke_hook( gate: _SessionHookGate | None, config: Mapping[str, Any] | None, deadline: timedelta, + logical_event_id: str | None = None, ) -> HookDecision: if gate is None: msg = "hooks_snapshot_id is required to emit server-owned hook events" @@ -544,6 +607,7 @@ def _invoke_hook( snapshot_id=gate["snapshot_id"], context=context, event=event, + logical_event_id=logical_event_id, ) request = HookInvocationRequest( protocol_version=1, @@ -632,17 +696,23 @@ def _invocation_id( event: ( PreToolUseEvent | PostToolUseEvent + | PreCompactEvent | StopEvent | SubagentStartEvent | SubagentStopEvent ), + logical_event_id: str | None = None, ) -> UUID: identity = { "run_id": run_id, "thread_id": context.thread_id, "snapshot_id": snapshot_id, "event": event.event.value, - "logical_event": _logical_event_identity(context, event), + "logical_event": _logical_event_identity( + context, + event, + logical_event_id=logical_event_id, + ), } return uuid5( _INVOCATION_NAMESPACE, @@ -655,13 +725,21 @@ def _logical_event_identity( event: ( PreToolUseEvent | PostToolUseEvent + | PreCompactEvent | StopEvent | SubagentStartEvent | SubagentStopEvent ), + *, + logical_event_id: str | None = None, ) -> str: if isinstance(event, PreToolUseEvent | PostToolUseEvent): return event.call.id + if isinstance(event, PreCompactEvent): + if logical_event_id: + return logical_event_id + msg = "PreCompact requires a stable tool-call identity" + raise ValueError(msg) if isinstance(event, SubagentStartEvent): return event.agent.id prompt_id = str(context.prompt_id) if context.prompt_id is not None else "" diff --git a/libs/code/deepagents_code/hooks/snapshot.py b/libs/code/deepagents_code/hooks/snapshot.py index fa6bdeb9b3..f5081b2a17 100644 --- a/libs/code/deepagents_code/hooks/snapshot.py +++ b/libs/code/deepagents_code/hooks/snapshot.py @@ -23,6 +23,7 @@ SubagentStartEvent, SubagentStopEvent, ) +from deepagents_code.hooks.projection import to_wire_notification_type from deepagents_code.hooks.tools import to_wire_tool_name if TYPE_CHECKING: @@ -241,7 +242,7 @@ def _match_target( ): return to_wire_tool_name(event.call.name, mcp_server=event.call.mcp_server) if matcher_field == "notification_type" and isinstance(event, NotificationEvent): - return event.notification.type + return to_wire_notification_type(event.notification.type).value if matcher_field == "cause" and isinstance( event, SessionStartEvent | SessionEndEvent ): diff --git a/libs/code/deepagents_code/hooks/transcript.py b/libs/code/deepagents_code/hooks/transcript.py index 7b0fddf062..2aec1d277d 100644 --- a/libs/code/deepagents_code/hooks/transcript.py +++ b/libs/code/deepagents_code/hooks/transcript.py @@ -29,9 +29,11 @@ from langchain_core.messages import ( AIMessage, BaseMessage, + BaseMessageChunk, HumanMessage, SystemMessage, ToolMessage, + message_chunk_to_message, ) from pydantic import BaseModel, ConfigDict @@ -39,10 +41,24 @@ from deepagents_code.json_types import JSON_VALUE_ADAPTER, JsonValue if TYPE_CHECKING: - from collections.abc import Sequence + from collections.abc import Mapping, Sequence + from typing import Protocol + + class _TranscriptRuntime(Protocol): + def append_messages( + self, + thread_id: str, + messages: Sequence[BaseMessage], + *, + agent_id: str | None = None, + ) -> None: ... + logger = logging.getLogger(__name__) +SUBAGENT_TRANSCRIPT_ID_METADATA_KEY = "dcode_subagent_id" +_INTERNAL_STREAM_SOURCES = frozenset({"summarization", "auto_mode_classifier"}) + TRANSCRIPT_SCHEMA_VERSION = 1 DEFAULT_RETENTION_REVISIONS = 20 _FILE_MODE = 0o600 @@ -120,6 +136,7 @@ class TranscriptHandle: @dataclass class _TranscriptBuffer: records: list[TranscriptRecord] = field(default_factory=list) + record_ids: set[str] = field(default_factory=set) dirty: bool = False revision: str = _EMPTY_REVISION @@ -205,7 +222,13 @@ def append_messages( ) if record is None: continue + if ( + record.message_id is not None + and record.record_id in buffer.record_ids + ): + continue buffer.records.append(record) + buffer.record_ids.add(record.record_id) buffer.dirty = True def materialize( @@ -277,12 +300,84 @@ def _buffer(self, thread_id: str, agent_id: str | None) -> _TranscriptBuffer: ) if path.is_file(): buffer.records, valid = _read_transcript(path) + buffer.record_ids = {record.record_id for record in buffer.records} buffer.revision = _revision_for_records(buffer.records) buffer.dirty = not valid self._buffers[key] = buffer return buffer +@dataclass(slots=True) +class TranscriptRecorder: + """Collect completed stream messages into a Hooks transcript runtime.""" + + runtime: _TranscriptRuntime + thread_id: str + _chunks: dict[tuple[str | None, str], BaseMessageChunk] = field( + default_factory=dict + ) + + def record( + self, + message: object, + metadata: Mapping[str, object] | None, + *, + main_agent: bool, + ) -> None: + """Record one streamed message when its transcript identity is stable. + + Args: + message: Streamed LangChain message or chunk. + metadata: Stream metadata carrying optional subagent identity. + main_agent: Whether the message belongs to the root graph. + """ + if ( + metadata is not None + and metadata.get("lc_source") in _INTERNAL_STREAM_SOURCES + ): + return + agent_id = None if main_agent else _stream_agent_id(metadata) + if not main_agent and agent_id is None: + return + if isinstance(message, BaseMessageChunk): + key = (agent_id, message.id or type(message).__name__) + previous = self._chunks.get(key) + combined = message if previous is None else previous + message + self._chunks[key] = combined + if getattr(message, "chunk_position", None) != "last": + return + self._chunks.pop(key, None) + self._append(message_chunk_to_message(combined), agent_id=agent_id) + return + if isinstance(message, BaseMessage): + self._append(message, agent_id=agent_id) + + def append(self, messages: Sequence[BaseMessage]) -> None: + """Append checkpoint or input messages to the root transcript.""" + append_messages = getattr(self.runtime, "append_messages", None) + if callable(append_messages): + append_messages(self.thread_id, messages) + + def _append(self, message: BaseMessage, *, agent_id: str | None) -> None: + append_messages = getattr(self.runtime, "append_messages", None) + if not callable(append_messages): + return + try: + append_messages(self.thread_id, [message], agent_id=agent_id) + except (TypeError, ValueError): + logger.warning( + "Skipping invalid streamed transcript message", + exc_info=True, + ) + + +def _stream_agent_id(metadata: Mapping[str, object] | None) -> str | None: + if metadata is None: + return None + value = metadata.get(SUBAGENT_TRANSCRIPT_ID_METADATA_KEY) + return value if isinstance(value, str) and value else None + + def _record_from_message( message: BaseMessage, *, diff --git a/libs/code/deepagents_code/tui/textual_adapter.py b/libs/code/deepagents_code/tui/textual_adapter.py index af5b09a6f6..934975731d 100644 --- a/libs/code/deepagents_code/tui/textual_adapter.py +++ b/libs/code/deepagents_code/tui/textual_adapter.py @@ -36,6 +36,7 @@ ClientHookService, ) from deepagents_code.hooks.models.domain import DcodeNotificationKind, HookEvent + from deepagents_code.hooks.runtime import HooksRuntime from deepagents_code.resume_state import RubricResult # Type alias matching HITLResponse["decisions"] element type @@ -56,6 +57,7 @@ class _ClientHookSessionState(Protocol): approval_mode: ApprovalMode turn_id: str | None client_hooks: ClientHookService | None + hooks_runtime: HooksRuntime | None from deepagents_code._ask_user_types import AskUserRequest @@ -87,6 +89,10 @@ class _ClientHookSessionState(Protocol): dispatch_hook, dispatch_hook_fire_and_forget, ) +from deepagents_code.hooks.client_lifecycle import ( + PermissionHookOutcome as _PermissionHookOutcome, + permission_hook_outcome, +) from deepagents_code.input import MediaTracker, parse_file_mentions from deepagents_code.media_utils import create_multimodal_content from deepagents_code.tool_display import format_tool_message_content @@ -107,11 +113,6 @@ class _ClientHookSessionState(Protocol): _ASK_USER_UNSUPPORTED_ERROR = "ask_user not supported by this UI" -class _PermissionHookOutcome(NamedTuple): - decision: dict[str, str] | None - interrupt: bool - - def _client_hook_context(session_state: _ClientHookSessionState) -> ClientHookContext: from deepagents_code.approval_mode import ApprovalMode, coerce_approval_mode from deepagents_code.hooks.client_lifecycle import ClientHookContext @@ -197,25 +198,7 @@ async def _permission_hook_outcomes( logger.warning("PermissionRequest hook invocation failed", exc_info=True) outcomes.append(_PermissionHookOutcome(None, False)) continue - if not hook_decision.continue_processing: - reason = hook_decision.stop_reason or "Permission stopped by hook" - outcomes.append( - _PermissionHookOutcome( - {"type": "reject", "message": reason}, - True, - ) - ) - continue - permission = hook_decision.permission - if permission.behavior == "allow": - outcomes.append(_PermissionHookOutcome({"type": "approve"}, False)) - elif permission.behavior == "deny": - decision = {"type": "reject"} - if permission.reason: - decision["message"] = permission.reason - outcomes.append(_PermissionHookOutcome(decision, permission.interrupt)) - else: - outcomes.append(_PermissionHookOutcome(None, False)) + outcomes.append(permission_hook_outcome(hook_decision)) return outcomes @@ -223,13 +206,27 @@ def _merge_permission_outcomes( outcomes: list[_PermissionHookOutcome], reviewed: Sequence[HITLDecision], ) -> list[HITLDecision]: + from langchain.agents.middleware.human_in_the_loop import ( + ApproveDecision, + RejectDecision, + ) + reviewed_iter = iter(reviewed) - return [ - cast("HITLDecision", outcome.decision) - if outcome.decision is not None - else next(reviewed_iter) - for outcome in outcomes - ] + merged: list[HITLDecision] = [] + for outcome in outcomes: + decision = outcome.decision + if decision is None: + merged.append(next(reviewed_iter)) + elif decision["type"] == "approve": + merged.append(ApproveDecision(type="approve")) + else: + message = decision.get("message") + merged.append( + RejectDecision(type="reject", message=message) + if message + else RejectDecision(type="reject") + ) + return merged def _dispatch_tool_use_hook( @@ -905,6 +902,7 @@ async def execute_task_textual( from deepagents_code.approval_mode import ApprovalMode, awrite_approval_mode from deepagents_code.auto_mode import USER_PROMPT_METADATA_KEY, user_prompt_metadata + from deepagents_code.hooks.models.domain import HookEvent hitl_request_adapter = _get_hitl_request_adapter(HITLRequest) ask_user_adapter = _get_ask_user_adapter() @@ -972,8 +970,6 @@ async def execute_task_textual( auto_approve=bool(session_state.auto_approve), ) - await dispatch_hook("session.start", {"thread_id": thread_id}) - captured_input_tokens = 0 captured_output_tokens = 0 if turn_stats is None: @@ -1040,6 +1036,12 @@ def _notify_user_visible_output_started() -> None: # when multiple subagents stream in parallel pending_text_by_namespace: dict[tuple, str] = {} assistant_message_by_namespace: dict[tuple, Any] = {} + hooks_runtime = getattr(session_state, "hooks_runtime", None) + transcript = None + if hooks_runtime is not None: + from deepagents_code.hooks.transcript import TranscriptRecorder + + transcript = TranscriptRecorder(hooks_runtime, thread_id) if image_tracker and graph_input is None: image_tracker.clear() @@ -1060,6 +1062,27 @@ def _notify_user_visible_output_started() -> None: user_msg["additional_kwargs"] = trusted_kwargs messages: list[dict[str, Any]] = [] client_hooks = getattr(session_state, "client_hooks", None) + if transcript is not None: + transcript.append([HumanMessage(content=message_content or "")]) + prompt_decision = None + if client_hooks is not None and client_hooks.has_handlers( + HookEvent.USER_PROMPT_SUBMIT + ): + prompt_decision = await client_hooks.user_prompt_submit( + _client_hook_context(session_state), + user_input, + ) + if not prompt_decision.continue_processing: + from deepagents_code.hooks.client_lifecycle import ClientHookStopError + + reason = ( + prompt_decision.stop_reason + or "User prompt submission stopped by hook" + ) + raise ClientHookStopError(reason) + else: + await dispatch_hook("session.start", {"thread_id": thread_id}) + await dispatch_hook("user.prompt", {}) if client_hooks is not None: session_context = client_hooks.take_session_context(thread_id) if session_context: @@ -1069,7 +1092,15 @@ def _notify_user_visible_output_started() -> None: "content": "\n\n".join(session_context), } ) - messages.append(user_msg) + if prompt_decision is not None and prompt_decision.context: + messages.append( + { + "role": "system", + "content": "\n\n".join(prompt_decision.context), + } + ) + if prompt_decision is None or not prompt_decision.suppress_original_prompt: + messages.append(user_msg) stream_input: dict | Command = { "messages": messages, "goal_criteria_request": None, @@ -1084,7 +1115,28 @@ def _notify_user_visible_output_started() -> None: # Track summarization lifecycle so spinner status and notification stay in sync. summarization_in_progress = False - summarization_observed = False + completed_compaction_ids: set[str] = set() + + async def _after_automatic_compact() -> None: + from deepagents_code.config import settings + from deepagents_code.hooks.client_lifecycle import ClientHookStopError + from deepagents_code.hooks.models.domain import SessionStartCause + + service = getattr(session_state, "client_hooks", None) + if service is None: + return + try: + decision = await service.session_start( + _client_hook_context(session_state), + SessionStartCause.COMPACT, + model=settings.model_name or None, + ) + except Exception: + logger.warning("Compact SessionStart hook invocation failed", exc_info=True) + return + if not decision.continue_processing: + reason = decision.stop_reason or "Compact session start stopped by hook" + raise ClientHookStopError(reason) try: while True: @@ -1431,11 +1483,6 @@ def _notify_user_visible_output_started() -> None: # Handle MESSAGES stream - for content and tool calls elif current_stream_mode == "messages": - # Skip subagent outputs - only render main agent content in chat - if not is_main_agent: - logger.debug("Skipping subagent message ns=%s", ns_key) - continue - if not isinstance(data, tuple) or len(data) != 2: # noqa: PLR2004 # message stream data is a 2-tuple (message, metadata) logger.debug( "Skipping non-2-tuple message data: type=%s", @@ -1444,6 +1491,16 @@ def _notify_user_visible_output_started() -> None: continue message, metadata = data + if transcript is not None: + transcript.record( + message, + metadata if isinstance(metadata, dict) else None, + main_agent=is_main_agent, + ) + # Skip subagent outputs - only render main agent content in chat + if not is_main_agent: + logger.debug("Skipping subagent message ns=%s", ns_key) + continue logger.debug( "Processing message: type=%s id=%s has_content_blocks=%s", type(message).__name__, @@ -1457,7 +1514,6 @@ def _notify_user_visible_output_started() -> None: # These are hidden from the user; only the spinner and a # notification widget provide feedback. if _is_summarization_chunk(metadata): - summarization_observed = True if not summarization_in_progress: summarization_in_progress = True if adapter._set_spinner: @@ -1507,6 +1563,17 @@ def _notify_user_visible_output_started() -> None: # has finished. Mount the notification and reset the spinner. if summarization_in_progress: summarization_in_progress = False + if isinstance(message, ToolMessage): + raw_id = getattr(message, "tool_call_id", None) + if ( + isinstance(raw_id, str) + and raw_id + and str(message.content).startswith( + "Conversation compacted." + ) + ): + completed_compaction_ids.add(raw_id) + await _after_automatic_compact() try: await adapter._mount_message(SummarizationMessage()) except Exception: @@ -1558,6 +1625,15 @@ def _notify_user_visible_output_started() -> None: except Exception: logger.exception("Failed to format tool output") output_str = UNRENDERABLE_TOOL_OUTPUT + compaction_id = getattr(message, "tool_call_id", None) + if ( + isinstance(compaction_id, str) + and compaction_id + and compaction_id not in completed_compaction_ids + and output_str.startswith("Conversation compacted.") + ): + completed_compaction_ids.add(compaction_id) + await _after_automatic_compact() record = file_op_tracker.complete_with_message(message) # Update tool call status with output @@ -1867,6 +1943,7 @@ def _notify_user_visible_output_started() -> None: # (e.g. middleware error, stream exhausted before regular chunks). if summarization_in_progress: summarization_in_progress = False + await _after_automatic_compact() try: await adapter._mount_message(SummarizationMessage()) except Exception: @@ -1876,31 +1953,6 @@ def _notify_user_visible_output_started() -> None: ) if adapter._set_spinner and not adapter._current_tool_messages: await adapter._set_spinner("Thinking") - if summarization_observed: - from deepagents_code.hooks.client_lifecycle import ClientHookStopError - from deepagents_code.hooks.models.domain import SessionStartCause - - service = getattr(session_state, "client_hooks", None) - if service is not None: - try: - decision = await service.session_start( - _client_hook_context(session_state), - SessionStartCause.COMPACT, - ) - except Exception: - logger.warning( - "Compact SessionStart hook invocation failed", - exc_info=True, - ) - else: - if not decision.continue_processing: - reason = ( - decision.stop_reason - or "Compact session start stopped by hook" - ) - raise ClientHookStopError(reason) - summarization_observed = False - # Flush any remaining text from all namespaces for ns_key, pending_text in list(pending_text_by_namespace.items()): if pending_text: @@ -2192,17 +2244,22 @@ def _notify_user_visible_output_started() -> None: adapter._current_tool_messages, ) if any(outcome.interrupt for outcome in hook_outcomes): - decisions = [ - cast( - "HITLDecision", - outcome.decision - or { + interrupted_outcomes = [ + outcome + if outcome.decision is not None + else _PermissionHookOutcome( + { "type": "reject", "message": "Permission interrupted by hook", }, + interrupt=True, ) for outcome in hook_outcomes ] + decisions = _merge_permission_outcomes( + interrupted_outcomes, + [], + ) for tool_msg in _interrupt_tool_rows( namespace, all_action_requests, diff --git a/libs/code/tests/unit_tests/hooks/test_client_lifecycle.py b/libs/code/tests/unit_tests/hooks/test_client_lifecycle.py index 25eb128f60..857210e6b2 100644 --- a/libs/code/tests/unit_tests/hooks/test_client_lifecycle.py +++ b/libs/code/tests/unit_tests/hooks/test_client_lifecycle.py @@ -20,6 +20,7 @@ ClientHookStopError, ) from deepagents_code.hooks.models.domain import ( + CompactTrigger, DcodeNotificationKind, HookDecision, HookDiagnostic, @@ -28,10 +29,12 @@ NotificationDecision, PermissionEffect, PermissionRequestDecision, + PreCompactDecision, SessionEndCause, SessionEndDecision, SessionStartCause, SessionStartDecision, + UserPromptSubmitDecision, ) from deepagents_code.hooks.models.wire import NotificationWireInput from deepagents_code.hooks.projection import project_hook_input @@ -49,6 +52,8 @@ RejectDecision, ) + from deepagents_code.hooks.runtime import HooksRuntime + HITLDecision = ApproveDecision | EditDecision | RejectDecision @@ -72,6 +77,7 @@ class _SessionState: approval_mode: ApprovalMode turn_id: str | None client_hooks: ClientHookService | None + hooks_runtime: HooksRuntime | None = None def _context() -> ClientHookContext: @@ -190,6 +196,34 @@ async def test_session_end_discards_pending_context(tmp_path: Path) -> None: ] +async def test_prompt_and_compact_services_preserve_typed_effects( + tmp_path: Path, +) -> None: + runtime = _Runtime( + cwd=tmp_path, + decisions=deque( + [ + UserPromptSubmitDecision( + event=HookEvent.USER_PROMPT_SUBMIT, + context=["prompt context"], + suppress_original_prompt=True, + ), + PreCompactDecision(event=HookEvent.PRE_COMPACT), + ] + ), + ) + service = ClientHookService(runtime) + + prompt = await service.user_prompt_submit(_context(), "hello") + compact = await service.pre_compact(_context(), CompactTrigger.AUTO) + + assert prompt.context == ["prompt context"] + assert prompt.suppress_original_prompt is True + assert compact.event is HookEvent.PRE_COMPACT + assert runtime.invocations[0].event.event is HookEvent.USER_PROMPT_SUBMIT + assert runtime.invocations[1].event.event is HookEvent.PRE_COMPACT + + async def test_notification_stop_interrupts_client_processing(tmp_path: Path) -> None: runtime = _Runtime( cwd=tmp_path, @@ -318,3 +352,54 @@ def _resolve(*_args: object) -> dict[str, str]: should_resolve = expected == {"type": "approve"} and len(runtime.invocations) == 2 assert resolution_calls == int(should_resolve) assert runtime.invocations[0].event.event is HookEvent.PERMISSION_REQUEST + + +async def test_headless_permission_uses_live_context(tmp_path: Path) -> None: + from uuid import uuid4 + + prompt_id = uuid4() + runtime = _Runtime(cwd=tmp_path, decisions=deque([_permission("allow")])) + context = ClientHookContext.create( + thread_id="thread-1", + approval_mode=ApprovalMode.AUTO, + prompt_id=prompt_id, + ) + state = StreamState( + client_hooks=ClientHookService(runtime), + client_hook_context=context, + ) + state.pending_interrupts["interrupt-1"] = { + "action_requests": [{"name": "read_file", "args": {"path": "README.md"}}], + "review_configs": [], + } + + await _process_hitl_interrupts(state, Console(quiet=True), "thread-1") + + invocation = runtime.invocations[0] + assert invocation.context.approval_mode is ApprovalMode.AUTO + assert invocation.context.prompt_id == prompt_id + + +async def test_headless_compact_permission_does_not_redispatch_precompact( + tmp_path: Path, +) -> None: + runtime = _Runtime( + cwd=tmp_path, + decisions=deque([_permission("allow")]), + ) + state = StreamState( + client_hooks=ClientHookService(runtime), + client_hook_context=_context(), + ) + state.pending_interrupts["interrupt-1"] = { + "action_requests": [ + {"name": "compact_conversation", "args": {}}, + ], + "review_configs": [], + } + + await _process_hitl_interrupts(state, Console(quiet=True), "thread-1") + + assert [invocation.event.event for invocation in runtime.invocations] == [ + HookEvent.PERMISSION_REQUEST, + ] diff --git a/libs/code/tests/unit_tests/hooks/test_configuration.py b/libs/code/tests/unit_tests/hooks/test_configuration.py index 49608fc1fd..4feafad637 100644 --- a/libs/code/tests/unit_tests/hooks/test_configuration.py +++ b/libs/code/tests/unit_tests/hooks/test_configuration.py @@ -24,7 +24,7 @@ ) from deepagents_code.hooks.migration import migrate_legacy_hooks from deepagents_code.hooks.models.config import HooksConfig -from deepagents_code.hooks.models.domain import HookEvent +from deepagents_code.hooks.models.domain import HookEvent, HookOwner from deepagents_code.hooks.snapshot import HooksSnapshot if TYPE_CHECKING: @@ -44,6 +44,7 @@ def test_registry_covers_all_hook_events() -> None: HookEvent.USER_PROMPT_SUBMIT ).default_timeout_seconds == pytest.approx(30.0) assert get_event_spec(HookEvent.PRE_COMPACT).matcher_field == "trigger" + assert get_event_spec(HookEvent.PRE_COMPACT).owner is HookOwner.SERVER def test_load_hooks_config_precedence_and_snapshot_hash(tmp_path: Path) -> None: @@ -164,6 +165,10 @@ def test_legacy_migration_maps_equivalent_lifecycle_events( ] assert prompt_legacy_events == ["session.start", "user.prompt"] assert compact_legacy_events == ["context.offload", "context.compact"] + assert ( + HookEvent.PRE_COMPACT + in HooksSnapshot.from_config(migrated).configured_server_events() + ) assert HookEvent.SESSION_START not in migrated.hooks assert HookEvent.PRE_TOOL_USE not in migrated.hooks for groups in migrated.hooks.values(): diff --git a/libs/code/tests/unit_tests/hooks/test_engine.py b/libs/code/tests/unit_tests/hooks/test_engine.py index c6abff4e28..f4cef15437 100644 --- a/libs/code/tests/unit_tests/hooks/test_engine.py +++ b/libs/code/tests/unit_tests/hooks/test_engine.py @@ -209,7 +209,7 @@ def test_snapshot_matches_notification_and_skips_tool_mismatch(tmp_path: Path) - { "Notification": [ { - "matcher": "permission_.*", + "matcher": "permission_prompt", "hooks": [{"type": "command", "command": "notify"}], } ], diff --git a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py index 24c967302c..01bea06fdb 100644 --- a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py +++ b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py @@ -34,11 +34,14 @@ from deepagents_code.hooks.models.adapters import HOOKS_CONFIG_ADAPTER from deepagents_code.hooks.models.config import HooksConfig from deepagents_code.hooks.models.domain import ( + CompactTrigger, HookContext, HookEvent, HookInvocation, PermissionEffect, PostToolUseDecision, + PreCompactDecision, + PreCompactEvent, PreToolUseDecision, PreToolUseEvent, StopDecision, @@ -63,6 +66,7 @@ _session_gate, ) from deepagents_code.hooks.snapshot import HooksSnapshot +from deepagents_code.hooks.transcript import SUBAGENT_TRANSCRIPT_ID_METADATA_KEY if TYPE_CHECKING: from langchain_core.runnables import RunnableConfig @@ -189,11 +193,69 @@ def invoke_hook(state: _ReplayState) -> dict[str, bool]: assert resumed["completed"] is True +def test_precompact_replay_uses_stable_tool_call_identity() -> None: + context = HookContext( + thread_id="thread-1", + cwd=Path("/tmp"), + approval_mode=ApprovalMode.MANUAL, + ) + event = PreCompactEvent( + event=HookEvent.PRE_COMPACT, + trigger=CompactTrigger.AUTO, + ) + gate = _session_gate( + { + "hooks_snapshot_id": "snapshot-1", + "hooks_server_events": [HookEvent.PRE_COMPACT.value], + } + ) + assert gate is not None + + def invoke_hook(state: _ReplayState) -> dict[str, bool]: + del state + decision = _invoke_hook( + context, + event, + gate=gate, + config={"configurable": {"thread_id": "thread-1"}}, + deadline=timedelta(minutes=1), + logical_event_id="compact-call-1", + ) + assert isinstance(decision, PreCompactDecision) + return {"completed": decision.continue_processing} + + builder = StateGraph(_ReplayState) + builder.add_node("hook", invoke_hook) + builder.add_edge(START, "hook") + graph = builder.compile(checkpointer=InMemorySaver()) + config: RunnableConfig = {"configurable": {"thread_id": "thread-1"}} + + interrupted = graph.invoke(_ReplayState(completed=False), config) + pending = interrupted["__interrupt__"][0] + request = parse_hook_interrupt_payload(pending.value) + assert request is not None + response = HookInvocationResponse( + protocol_version=1, + invocation_id=request.invocation_id, + snapshot_id=request.snapshot_id, + decision=PreCompactDecision(event=HookEvent.PRE_COMPACT), + ) + + resumed = graph.invoke(Command(resume=build_hook_resume_value(response)), config) + + assert resumed["completed"] is True + + def test_apply_hooks_context_sets_server_events(tmp_path: Path) -> None: config_dir = tmp_path / "config" config_dir.mkdir() (config_dir / "hooks.json").write_text( - '{"hooks":{"PreToolUse":[{"hooks":[{"type":"command","command":"true"}]}]}}', + ( + '{"hooks":{' + '"PreCompact":[{"hooks":[{"type":"command","command":"true"}]}],' + '"PreToolUse":[{"hooks":[{"type":"command","command":"true"}]}]' + "}}" + ), encoding="utf-8", ) runtime = HooksRuntime.create( @@ -205,9 +267,9 @@ def test_apply_hooks_context_sets_server_events(tmp_path: Path) -> None: apply_hooks_context(context, runtime, prompt_id="prompt-1") assert context["hooks_snapshot_id"] == runtime.snapshot_id - assert context["hooks_server_events"] == ["PreToolUse"] + assert context["hooks_server_events"] == ["PreCompact", "PreToolUse"] assert context["prompt_id"] == "prompt-1" - assert runtime.configured_server_events() == ("PreToolUse",) + assert runtime.configured_server_events() == ("PreCompact", "PreToolUse") def test_session_gate_requires_snapshot_and_events() -> None: @@ -531,6 +593,195 @@ def test_pre_tool_deny_skips_hitl_and_execution( handler.assert_not_called() +@pytest.mark.parametrize( + ("args", "expected"), + [ + ({"force": True}, CompactTrigger.MANUAL), + ({}, CompactTrigger.AUTO), + ], +) +def test_precompact_trigger_uses_forced_tool_identity( + monkeypatch: pytest.MonkeyPatch, + args: dict[str, object], + expected: CompactTrigger, +) -> None: + middleware = ServerHooksMiddleware(cwd=Path("/tmp")) + runtime = MagicMock() + runtime.context = { + "hooks_snapshot_id": "snap", + "hooks_server_events": ["PreCompact"], + "thread_id": "t1", + "approval_mode": "manual", + } + state: ServerHooksState = { + "messages": [ + AIMessage( + content="", + tool_calls=[ + { + "name": "compact_conversation", + "args": args, + "id": "compact-call", + "type": "tool_call", + } + ], + ) + ] + } + observed: list[tuple[CompactTrigger, str | None]] = [] + + def invoke( + _context: object, + event: object, + **kwargs: object, + ) -> PreCompactDecision: + assert isinstance(event, PreCompactEvent) + logical_id = kwargs.get("logical_event_id") + observed.append( + ( + event.trigger, + logical_id if isinstance(logical_id, str) else None, + ) + ) + return PreCompactDecision(event=HookEvent.PRE_COMPACT) + + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware._invoke_hook", + invoke, + ) + + update = middleware._after_model(state, runtime) + + assert observed == [(expected, "compact-call")] + assert update["_hooks_pre_tool_outcomes"]["compact-call"]["behavior"] == "none" + + +def test_precompact_runs_before_pretool( + monkeypatch: pytest.MonkeyPatch, +) -> None: + middleware = ServerHooksMiddleware(cwd=Path("/tmp")) + runtime = MagicMock() + runtime.context = { + "hooks_snapshot_id": "snap", + "hooks_server_events": ["PreCompact", "PreToolUse"], + "thread_id": "t1", + "approval_mode": "manual", + } + state: ServerHooksState = { + "messages": [ + AIMessage( + content="", + tool_calls=[ + { + "name": "compact_conversation", + "args": {}, + "id": "compact-call", + "type": "tool_call", + } + ], + ) + ] + } + order: list[HookEvent] = [] + + def invoke( + _context: object, + event: object, + **_kwargs: object, + ) -> PreCompactDecision | PreToolUseDecision: + assert isinstance(event, PreCompactEvent | PreToolUseEvent) + order.append(event.event) + if isinstance(event, PreCompactEvent): + return PreCompactDecision(event=HookEvent.PRE_COMPACT) + return PreToolUseDecision( + event=HookEvent.PRE_TOOL_USE, + permission=PermissionEffect(behavior="none"), + context=["continue with compacted context"], + ) + + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware._invoke_hook", + invoke, + ) + + update = middleware._after_model(state, runtime) + + assert order == [HookEvent.PRE_COMPACT, HookEvent.PRE_TOOL_USE] + assert update["_hooks_pre_tool_outcomes"]["compact-call"]["context"] == [ + "continue with compacted context" + ] + state["_hooks_pre_tool_outcomes"] = update["_hooks_pre_tool_outcomes"] + runtime.config = {"configurable": {"thread_id": "t1"}} + request = MagicMock() + request.state = state + request.runtime = runtime + request.tool = None + message = state["messages"][-1] + assert isinstance(message, AIMessage) + request.tool_call = message.tool_calls[0] + result = middleware.wrap_tool_call( + request, + lambda _request: ToolMessage( + content="compacted", + name="compact_conversation", + tool_call_id="compact-call", + ), + ) + assert isinstance(result, ToolMessage) + assert "continue with compacted context" in str(result.content) + + +def test_precompact_block_skips_hitl_and_tool_execution( + monkeypatch: pytest.MonkeyPatch, +) -> None: + middleware = ServerHooksMiddleware(cwd=Path("/tmp")) + runtime = MagicMock() + runtime.context = { + "hooks_snapshot_id": "snap", + "hooks_server_events": ["PreCompact", "PreToolUse"], + "thread_id": "t1", + "approval_mode": "manual", + } + runtime.config = {"configurable": {"thread_id": "t1"}} + tool_call = { + "name": "compact_conversation", + "args": {}, + "id": "compact-call", + "type": "tool_call", + } + state: ServerHooksState = { + "messages": [AIMessage(content="", tool_calls=[tool_call])] + } + invoke = MagicMock( + return_value=PreCompactDecision( + event=HookEvent.PRE_COMPACT, + continue_processing=False, + stop_reason="keep full context", + ) + ) + monkeypatch.setattr( + "deepagents_code.hooks.server_middleware._invoke_hook", + invoke, + ) + update = middleware._after_model(state, runtime) + state["_hooks_pre_tool_outcomes"] = update["_hooks_pre_tool_outcomes"] + request = MagicMock() + request.state = state + request.runtime = runtime + request.tool = None + request.tool_call = tool_call + handler = MagicMock() + + assert _should_interrupt_tool_call(request) is False + result = middleware.wrap_tool_call(request, handler) + + assert isinstance(result, ToolMessage) + assert result.status == "error" + assert "keep full context" in str(result.content) + invoke.assert_called_once() + handler.assert_not_called() + + def test_ask_permission_via_hitl_approve(monkeypatch: pytest.MonkeyPatch) -> None: call = ToolCallData(id="c1", name="execute", args={"command": "ls"}) @@ -648,6 +899,79 @@ def test_subagent_start_deny_returns_error_tool_message( handler.assert_not_called() +def test_task_tool_scopes_subagent_transcript_identity() -> None: + from langchain_core.runnables.config import ensure_config + + middleware = ServerHooksMiddleware(cwd=Path("/tmp")) + request = MagicMock() + request.state = {} + request.tool = None + request.tool_call = { + "name": "task", + "args": {"subagent_type": "researcher", "description": "go"}, + "id": "call-1", + "type": "tool_call", + } + request.runtime.context = { + "thread_id": "t1", + "approval_mode": "manual", + } + request.runtime.config = { + "configurable": {"thread_id": "t1"}, + "metadata": {"parent": "kept"}, + } + observed: list[dict[str, object]] = [] + + def handler(_request: object) -> ToolMessage: + nested = ensure_config({"configurable": {"ls_agent_type": "subagent"}}) + observed.append(dict(nested["metadata"])) + return ToolMessage(content="done", name="task", tool_call_id="call-1") + + result = middleware.wrap_tool_call(request, handler) + + assert isinstance(result, ToolMessage) + assert observed == [ + { + "parent": "kept", + SUBAGENT_TRANSCRIPT_ID_METADATA_KEY: "call-1", + } + ] + assert SUBAGENT_TRANSCRIPT_ID_METADATA_KEY not in ensure_config()["metadata"] + + +async def test_async_task_tool_scopes_subagent_transcript_identity() -> None: + from langchain_core.runnables.config import ensure_config + + middleware = ServerHooksMiddleware(cwd=Path("/tmp")) + request = MagicMock() + request.state = {} + request.tool = None + request.tool_call = { + "name": "task", + "args": {"subagent_type": "researcher", "description": "go"}, + "id": "call-1", + "type": "tool_call", + } + request.runtime.context = { + "thread_id": "t1", + "approval_mode": "manual", + } + request.runtime.config = {"configurable": {"thread_id": "t1"}} + observed: list[str | None] = [] + + async def handler(_request: object) -> ToolMessage: + await asyncio.sleep(0) + nested = ensure_config({"configurable": {"ls_agent_type": "subagent"}}) + observed.append(nested["metadata"].get(SUBAGENT_TRANSCRIPT_ID_METADATA_KEY)) + return ToolMessage(content="done", name="task", tool_call_id="call-1") + + result = await middleware.awrap_tool_call(request, handler) + + assert isinstance(result, ToolMessage) + assert observed == ["call-1"] + assert SUBAGENT_TRANSCRIPT_ID_METADATA_KEY not in ensure_config()["metadata"] + + async def test_fulfill_hook_invocation_runs_engine(tmp_path: Path) -> None: config_dir = tmp_path / "config" config_dir.mkdir() @@ -728,6 +1052,9 @@ def test_snapshot_configured_server_events() -> None: "SessionStart": [ {"hooks": [{"type": "command", "command": "echo client"}]} ], + "PreCompact": [ + {"hooks": [{"type": "command", "command": "echo compact"}]} + ], "PreToolUse": [ {"hooks": [{"type": "command", "command": "echo server"}]} ], @@ -738,6 +1065,10 @@ def test_snapshot_configured_server_events() -> None: snapshot = HooksSnapshot.from_config(config) assert snapshot.configured_events() == { HookEvent.SESSION_START, + HookEvent.PRE_COMPACT, + HookEvent.PRE_TOOL_USE, + } + assert snapshot.configured_server_events() == { + HookEvent.PRE_COMPACT, HookEvent.PRE_TOOL_USE, } - assert snapshot.configured_server_events() == {HookEvent.PRE_TOOL_USE} diff --git a/libs/code/tests/unit_tests/hooks/test_transcript.py b/libs/code/tests/unit_tests/hooks/test_transcript.py index 7c130c3839..db37240a20 100644 --- a/libs/code/tests/unit_tests/hooks/test_transcript.py +++ b/libs/code/tests/unit_tests/hooks/test_transcript.py @@ -10,7 +10,7 @@ from typing import TYPE_CHECKING import pytest -from langchain_core.messages import AIMessage, HumanMessage +from langchain_core.messages import AIMessage, AIMessageChunk, HumanMessage from deepagents_code.approval_mode import ApprovalMode from deepagents_code.hooks.models.domain import ( @@ -24,7 +24,12 @@ SubagentStopEvent, ) from deepagents_code.hooks.runtime import HooksRuntime -from deepagents_code.hooks.transcript import TranscriptStore, redact_transcript_value +from deepagents_code.hooks.transcript import ( + SUBAGENT_TRANSCRIPT_ID_METADATA_KEY, + TranscriptRecorder, + TranscriptStore, + redact_transcript_value, +) if TYPE_CHECKING: from pathlib import Path @@ -183,6 +188,71 @@ def append(index: int) -> None: assert handle.revision == concurrent.revision("thread") +def test_transcript_deduplicates_stable_message_identity(tmp_path: Path) -> None: + store = TranscriptStore(tmp_path / "transcripts") + message = HumanMessage(id="user-1", content="hello") + + store.append_messages("thread", [message, message]) + store.append_messages("thread", [message]) + + records = store.materialize("thread").path.read_text(encoding="utf-8").splitlines() + assert len(records) == 1 + + +def test_stream_recorder_collects_completed_main_and_identified_subagent( + tmp_path: Path, +) -> None: + runtime = HooksRuntime.create( + cwd=tmp_path, + config_dir=tmp_path / "config", + transcript_root=tmp_path / "transcripts", + ) + recorder = TranscriptRecorder(runtime, "thread") + recorder.record( + AIMessageChunk(id="main-1", content="hel"), + {}, + main_agent=True, + ) + recorder.record( + AIMessageChunk(id="main-1", content="lo", chunk_position="last"), + {}, + main_agent=True, + ) + recorder.record( + AIMessage(id="sub-1", content="research"), + {SUBAGENT_TRANSCRIPT_ID_METADATA_KEY: "agent-1"}, + main_agent=False, + ) + recorder.record( + AIMessage(id="unstable", content="skip"), + {}, + main_agent=False, + ) + recorder.record( + AIMessage(id="summary", content="hidden summary"), + {"lc_source": "summarization"}, + main_agent=True, + ) + recorder.record( + AIMessage(id="classifier", content="hidden classifier"), + {"lc_source": "auto_mode_classifier"}, + main_agent=True, + ) + + main = runtime.transcripts.materialize("thread").path.read_text(encoding="utf-8") + agent = runtime.transcripts.materialize( + "thread", + agent_id="agent-1", + ).path.read_text(encoding="utf-8") + + assert '"content":"hello"' in main + assert '"content":"research"' in agent + assert "skip" not in main + assert "skip" not in agent + assert "hidden summary" not in main + assert "hidden classifier" not in main + + def test_runtime_stores_transcripts_outside_workspace( tmp_path: Path, monkeypatch: pytest.MonkeyPatch ) -> None: diff --git a/libs/code/tests/unit_tests/test_app.py b/libs/code/tests/unit_tests/test_app.py index 55bb362086..6724a1b7d2 100644 --- a/libs/code/tests/unit_tests/test_app.py +++ b/libs/code/tests/unit_tests/test_app.py @@ -581,15 +581,22 @@ async def test_stopped_session_start_blocks_initial_submission(self) -> None: initial_prompt="hello", ) app._session_state = TextualSessionState(thread_id="thread-123") - app._run_session_start_hook = AsyncMock( # ty: ignore[invalid-assignment] - return_value=False - ) + hook = AsyncMock(return_value=False) + app._run_session_start_hook = hook # ty: ignore[invalid-assignment] submit = AsyncMock() + drain = AsyncMock() app._submit_initial_submission = submit # ty: ignore[invalid-assignment] + app._drain_startup_backlog = drain # ty: ignore[invalid-assignment] + await app._run_session_start_sequence() + await app._run_session_start_sequence() await app._run_session_start_sequence() submit.assert_not_awaited() + drain.assert_not_awaited() + hook.assert_awaited_once() + assert app._initial_session_started is False + assert app._initial_session_start_stopped is True async def test_reconnect_drains_queue_without_reloading_history(self) -> None: """Later `ServerReady` events should drain queued input once connected.""" @@ -640,16 +647,21 @@ async def capture_startup(command: str) -> None: # noqa: RUF029 assert command == "echo hi" order.append("startup") + async def capture_hook(_cause: object) -> bool: # noqa: RUF029 + order.append("hook") + return True + async def capture_goal_review() -> None: # noqa: RUF029 order.append("goal") app._load_thread_history = capture_history # ty: ignore + app._run_session_start_hook = capture_hook # ty: ignore[invalid-assignment] app._run_startup_command = capture_startup # ty: ignore app._remount_pending_goal_rubric_review = capture_goal_review # ty: ignore await app._run_session_start_sequence() - assert order == ["history", "startup", "goal"] + assert order == ["history", "hook", "startup", "goal"] assert app._startup_sequence_running is False @pytest.mark.parametrize( @@ -12962,6 +12974,34 @@ async def test_footers_render_for_restored_thread_history( with pytest.raises(NoMatches): app.query_one("#hist-app-timestamp-footer", Static) + async def test_resumed_history_populates_hook_transcript(self) -> None: + from langchain_core.messages import HumanMessage + + from deepagents_code.app import _ThreadHistoryPayload + + app = DeepAgentsApp() + async with app.run_test() as pilot: + await pilot.pause() + assert app._session_state is not None + runtime = MagicMock() + app._session_state.hooks_runtime = runtime + payload = _ThreadHistoryPayload( + messages=[], + context_tokens=0, + model_spec="", + transcript_messages=(HumanMessage(id="history-1", content="restored"),), + ) + + await app._load_thread_history( + thread_id="t-restored", + preloaded_payload=payload, + ) + + runtime.append_messages.assert_called_once_with( + "t-restored", + payload.transcript_messages, + ) + async def test_load_thread_history_skips_duplicate_ids(self) -> None: """History reusing an already-mounted widget ID is skipped, not fatal. @@ -25712,16 +25752,14 @@ async def test_toggle_off_while_reconnecting_stages_manual(self) -> None: async def test_session_init_keeps_mode_changed_during_construction(self) -> None: from deepagents_code.approval_mode import ApprovalMode - app = DeepAgentsApp() + app = DeepAgentsApp(approval_mode=ApprovalMode.MANUAL) - async def create_stale_session_state(*_args: object) -> TextualSessionState: - await asyncio.sleep(0) + def change_mode_during_construction(**_kwargs: object) -> None: app._approval_mode = ApprovalMode.AUTO - return TextualSessionState(approval_mode=ApprovalMode.MANUAL) with patch( - "deepagents_code.app.asyncio.to_thread", - new=create_stale_session_state, + "deepagents_code.hooks.runtime.HooksRuntime.create", + side_effect=change_mode_during_construction, ): await app._init_session_state() diff --git a/libs/code/tests/unit_tests/test_non_interactive.py b/libs/code/tests/unit_tests/test_non_interactive.py index 67eed2381d..b58054c0ba 100644 --- a/libs/code/tests/unit_tests/test_non_interactive.py +++ b/libs/code/tests/unit_tests/test_non_interactive.py @@ -41,11 +41,18 @@ _run_agent_loop, _run_startup_command, _start_langsmith_thread_url_lookup, + _summarization_stream_status, run_non_interactive, ) from deepagents_code.config import SHELL_ALLOW_ALL, ModelResult from deepagents_code.file_ops import FileOpTracker -from deepagents_code.hooks.models.domain import HookEvent +from deepagents_code.hooks.client_lifecycle import ClientHookStopError +from deepagents_code.hooks.models.domain import ( + HookEvent, + SessionEndDecision, + SessionStartDecision, + UserPromptSubmitDecision, +) from deepagents_code.tool_display import format_tool_message_content @@ -55,6 +62,32 @@ def console() -> Console: return Console(quiet=True) +def test_summarization_status_ignores_subagent_namespaces() -> None: + assert ( + _summarization_stream_status( + ( + ("tools:subagent",), + "messages", + ( + AIMessage(content="summary"), + {"lc_source": "summarization"}, + ), + ) + ) + is None + ) + assert ( + _summarization_stream_status( + ( + ("tools:subagent",), + "messages", + (AIMessage(content="continued"), {}), + ) + ) + is None + ) + + @pytest.fixture(autouse=True) def skip_mcp_metadata_preload() -> Iterator[None]: """Keep non-MCP non-interactive tests from starting connector discovery.""" @@ -1368,6 +1401,208 @@ async def test_run_agent_loop_passes_thread_id_context(self) -> None: _, kwargs = agent.astream.call_args assert kwargs["context"]["thread_id"] == "t1" + async def test_user_prompt_hook_suppresses_legacy_duplicate_and_prompt( + self, + tmp_path: Path, + ) -> None: + runtime = MagicMock() + runtime.cwd = tmp_path + runtime.snapshot_id = "snapshot" + runtime.configured_server_events.return_value = () + runtime.configured_events.return_value = frozenset( + {HookEvent.USER_PROMPT_SUBMIT} + ) + runtime.invoke = AsyncMock( + return_value=UserPromptSubmitDecision( + event=HookEvent.USER_PROMPT_SUBMIT, + context=["replacement"], + suppress_original_prompt=True, + ) + ) + agent = MagicMock() + agent.astream = MagicMock(return_value=_async_iter([])) + config: RunnableConfig = {"configurable": {"thread_id": "t1"}} + + with patch( + "deepagents_code.client.non_interactive.dispatch_hook", + new_callable=AsyncMock, + ) as legacy: + await _run_agent_loop( + agent, + "secret", + config, + Console(quiet=True), + MagicMock(), + quiet=True, + hooks_runtime=runtime, + ) + + stream_input = agent.astream.call_args.args[0] + assert stream_input["messages"] == [ + {"role": "system", "content": "replacement"} + ] + assert not any( + call.args and call.args[0] in {"session.start", "user.prompt"} + for call in legacy.await_args_list + ) + runtime.append_messages.assert_called_once() + + async def test_user_prompt_stop_ends_headless_session_once( + self, + tmp_path: Path, + ) -> None: + runtime = MagicMock() + runtime.cwd = tmp_path + runtime.snapshot_id = "snapshot" + runtime.configured_server_events.return_value = () + runtime.configured_events.return_value = frozenset( + { + HookEvent.SESSION_START, + HookEvent.USER_PROMPT_SUBMIT, + HookEvent.SESSION_END, + } + ) + runtime.invoke = AsyncMock( + side_effect=[ + SessionStartDecision(event=HookEvent.SESSION_START), + UserPromptSubmitDecision( + event=HookEvent.USER_PROMPT_SUBMIT, + continue_processing=False, + stop_reason="blocked", + ), + SessionEndDecision(event=HookEvent.SESSION_END), + ] + ) + agent = MagicMock() + + with pytest.raises(ClientHookStopError, match="blocked"): + await _run_agent_loop( + agent, + "secret", + {"configurable": {"thread_id": "t1"}}, + Console(quiet=True), + MagicMock(), + quiet=True, + hooks_runtime=runtime, + ) + + events = [call.args[0].event.event for call in runtime.invoke.await_args_list] + assert events == [ + HookEvent.SESSION_START, + HookEvent.USER_PROMPT_SUBMIT, + HookEvent.SESSION_END, + ] + agent.astream.assert_not_called() + + async def test_headless_transcript_records_main_and_identified_subagent( + self, + tmp_path: Path, + ) -> None: + runtime = MagicMock() + runtime.cwd = tmp_path + runtime.snapshot_id = "snapshot" + runtime.configured_server_events.return_value = () + runtime.configured_events.return_value = frozenset() + chunks = [ + ( + ("subagent",), + "messages", + ( + AIMessage(id="sub-1", content="research"), + {"dcode_subagent_id": "agent-1"}, + ), + ), + ( + (), + "messages", + (AIMessage(id="main-1", content="answer"), {}), + ), + ] + agent = MagicMock() + agent.astream = MagicMock(return_value=_async_iter(chunks)) + + with patch( + "deepagents_code.client.non_interactive.dispatch_hook", + new_callable=AsyncMock, + ): + await _run_agent_loop( + agent, + "question", + {"configurable": {"thread_id": "t1"}}, + Console(quiet=True), + MagicMock(), + quiet=True, + hooks_runtime=runtime, + ) + + calls = runtime.append_messages.call_args_list + assert calls[0].args[1][0].content == "question" + assert calls[1].kwargs == {"agent_id": "agent-1"} + assert calls[1].args[1][0].content == "research" + assert calls[2].kwargs == {"agent_id": None} + assert calls[2].args[1][0].content == "answer" + + async def test_compact_session_start_uses_active_model_before_continuation( + self, + tmp_path: Path, + ) -> None: + runtime = MagicMock() + runtime.cwd = tmp_path + runtime.snapshot_id = "snapshot" + runtime.configured_server_events.return_value = () + runtime.configured_events.return_value = frozenset( + {HookEvent.SESSION_START, HookEvent.PRE_COMPACT} + ) + runtime.invoke = AsyncMock( + side_effect=[ + SessionStartDecision(event=HookEvent.SESSION_START), + SessionStartDecision(event=HookEvent.SESSION_START), + ] + ) + chunks = [ + ( + (), + "messages", + ( + AIMessage(id="summary", content="summary"), + {"lc_source": "summarization"}, + ), + ), + ( + (), + "messages", + (AIMessage(id="answer", content="continued"), {}), + ), + ] + agent = MagicMock() + agent.astream = MagicMock(return_value=_async_iter(chunks)) + + with ( + patch( + "deepagents_code.client.non_interactive.dispatch_hook", + new_callable=AsyncMock, + ), + patch("deepagents_code.client.non_interactive.settings") as mock_settings, + ): + mock_settings.model_name = "test:model" + mock_settings.model_provider = "test" + await _run_agent_loop( + agent, + "question", + {"configurable": {"thread_id": "t1"}}, + Console(quiet=True), + MagicMock(), + quiet=True, + hooks_runtime=runtime, + ) + + invocations = [call.args[0] for call in runtime.invoke.await_args_list] + assert [invocation.event.event for invocation in invocations] == [ + HookEvent.SESSION_START, + HookEvent.SESSION_START, + ] + assert invocations[-1].event.model == "test:model" + async def test_run_agent_loop_defaults_project_hooks_untrusted( self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch ) -> None: diff --git a/libs/code/tests/unit_tests/test_offload.py b/libs/code/tests/unit_tests/test_offload.py index 98f2bb4f97..23381aabe6 100644 --- a/libs/code/tests/unit_tests/test_offload.py +++ b/libs/code/tests/unit_tests/test_offload.py @@ -1699,24 +1699,122 @@ async def test_seeds_tool_call_and_resumes_interrupt(self) -> None: assert isinstance(astream_inputs[1], Command) resume = astream_inputs[1].resume assert "interrupt-1" in resume - assert astream_contexts == [ - { - "model": "provider:active-model", - "model_params": {"temperature": 0}, - "profile_overrides": {"max_input_tokens": 4096}, - "model_context_limit": 4096, - "thread_id": "test-thread", - "offload_tool_call_id": tool_call["id"], - }, - { - "model": "provider:active-model", - "model_params": {"temperature": 0}, - "profile_overrides": {"max_input_tokens": 4096}, - "model_context_limit": 4096, - "thread_id": "test-thread", - "offload_tool_call_id": tool_call["id"], - }, - ] + expected = { + "model": "provider:active-model", + "model_params": {"temperature": 0}, + "profile_overrides": {"max_input_tokens": 4096}, + "model_context_limit": 4096, + "thread_id": "test-thread", + "offload_tool_call_id": tool_call["id"], + } + assert len(astream_contexts) == 2 + for context in astream_contexts: + assert isinstance(context, dict) + normalized = {str(key): value for key, value in context.items()} + assert {key: normalized[key] for key in expected} == expected + assert isinstance(normalized["hooks_snapshot_id"], str) + assert normalized["hooks_server_events"] == [] + + async def test_fulfills_precompact_before_manual_approval(self) -> None: + import asyncio + from types import SimpleNamespace + + from langchain_core.messages import ToolMessage + from langgraph.types import Command + + from deepagents_code.client.remote_client import RemoteAgent + from deepagents_code.hooks.interrupt import HOOK_INVOCATION_INTERRUPT_TYPE + + astream_inputs: list[Any] = [] + contexts: list[object] = [] + + async def _astream(stream_input: object, **kwargs: object): # noqa: ANN202 + await asyncio.sleep(0) + index = len(astream_inputs) + astream_inputs.append(stream_input) + contexts.append(kwargs.get("context")) + if index == 0: + yield ( + (), + "updates", + { + "__interrupt__": [ + SimpleNamespace( + id="hook-interrupt", + value={"type": HOOK_INVOCATION_INTERRUPT_TYPE}, + ) + ] + }, + ) + elif index == 1: + yield ( + (), + "updates", + { + "__interrupt__": [ + SimpleNamespace( + id="approval-interrupt", + value={ + "action_requests": [ + { + "name": "compact_conversation", + "args": {"force": True}, + } + ] + }, + ) + ] + }, + ) + else: + yield ( + (), + "messages", + ( + ToolMessage( + content="Conversation compacted. Summarized 2 messages.", + name="compact_conversation", + tool_call_id="compact-call", + ), + {}, + ), + ) + + agent = MagicMock(spec=RemoteAgent) + agent.aensure_thread = AsyncMock() + agent.aupdate_state = AsyncMock() + agent.astream = _astream + app = DeepAgentsApp() + async with app.run_test() as pilot: + await pilot.pause() + assert app._session_state is not None + runtime = MagicMock() + runtime.snapshot_id = "snapshot" + runtime.configured_server_events.return_value = ("PreCompact",) + app._session_state.hooks_runtime = runtime + app._agent = agent + app._lc_thread_id = "test-thread" + fulfill = AsyncMock(return_value={"hook": "approved"}) + + with patch( + "deepagents_code.hooks.client.fulfill_hook_interrupt", + fulfill, + ): + result = await app._drive_server_side_compaction( # ty: ignore + {"configurable": {"thread_id": "test-thread"}} + ) + + assert result is None + fulfill.assert_awaited_once() + assert len(astream_inputs) == 3 + assert isinstance(astream_inputs[1], Command) + assert astream_inputs[1].resume == {"hook-interrupt": {"hook": "approved"}} + assert isinstance(astream_inputs[2], Command) + assert "approval-interrupt" in astream_inputs[2].resume + for context in contexts: + assert isinstance(context, dict) + normalized = {str(key): value for key, value in context.items()} + assert normalized["hooks_server_events"] == ["PreCompact"] async def test_reports_tool_failure(self) -> None: """Returns the tool's error text when compaction fails.""" @@ -1766,18 +1864,20 @@ async def test_forwards_startup_model_profile_to_compaction(self) -> None: seed_values = agent.aupdate_state.call_args.args[1] (seed_msg,) = seed_values["messages"] (tool_call,) = seed_msg.tool_calls - assert all( - context - == { - "model": "provider:startup-model", - "model_params": {}, - "profile_overrides": {"max_input_tokens": 4096}, - "model_context_limit": 4096, - "thread_id": "test-thread", - "offload_tool_call_id": tool_call["id"], - } - for context in contexts - ) + expected = { + "model": "provider:startup-model", + "model_params": {}, + "profile_overrides": {"max_input_tokens": 4096}, + "model_context_limit": 4096, + "thread_id": "test-thread", + "offload_tool_call_id": tool_call["id"], + } + for context in contexts: + assert isinstance(context, dict) + normalized = {str(key): value for key, value in context.items()} + assert {key: normalized[key] for key in expected} == expected + assert isinstance(normalized["hooks_snapshot_id"], str) + assert normalized["hooks_server_events"] == [] async def test_rejects_interrupt_without_identifiable_action(self) -> None: """Malformed interrupt payloads fail closed instead of being approved.""" diff --git a/libs/code/tests/unit_tests/tui/test_textual_adapter.py b/libs/code/tests/unit_tests/tui/test_textual_adapter.py index e92d11255a..32c668859d 100644 --- a/libs/code/tests/unit_tests/tui/test_textual_adapter.py +++ b/libs/code/tests/unit_tests/tui/test_textual_adapter.py @@ -39,6 +39,8 @@ HookEvent, PermissionEffect, PermissionRequestDecision, + SessionStartDecision, + UserPromptSubmitDecision, ) from deepagents_code.tui.textual_adapter import ( RubricEvaluationEnd, @@ -56,6 +58,7 @@ ) from deepagents_code.tui.widgets.messages import ( AppMessage, + AssistantMessage, RubricResultMessage, SummarizationMessage, ToolCallMessage, @@ -1815,6 +1818,257 @@ async def test_auto_approve_absent_from_stream_config_when_disabled(self) -> Non assert "dcode_auto_approve" not in agent.configs[0]["metadata"] +class TestExecuteTaskTextualClientLifecycle: + async def test_transcript_records_completed_main_and_subagent_messages( + self, + ) -> None: + from deepagents_code.app import TextualSessionState + + runtime = MagicMock() + agent = _SequencedAgent( + [ + [ + ( + ("subagent",), + "messages", + ( + AIMessage(id="sub-1", content="research"), + {"dcode_subagent_id": "agent-1"}, + ), + ), + ( + (), + "messages", + (AIMessage(id="main-1", content="answer"), {}), + ), + ] + ] + ) + state = TextualSessionState(thread_id="thread-1") + state.hooks_runtime = runtime + adapter = TextualUIAdapter( + mount_message=_mock_mount, + update_status=_noop_status, + request_approval=_mock_approval, + ) + + await execute_task_textual( + user_input="question", + agent=agent, + assistant_id="assistant", + session_state=state, + adapter=adapter, + ) + + calls = runtime.append_messages.call_args_list + assert calls[0].args[1][0].content == "question" + assert calls[1].kwargs == {"agent_id": "agent-1"} + assert calls[1].args[1][0].content == "research" + assert calls[2].kwargs == {"agent_id": None} + assert calls[2].args[1][0].content == "answer" + + async def test_user_prompt_hook_applies_context_and_suppression_once(self) -> None: + from deepagents_code.app import TextualSessionState + + agent = _SequencedAgent([[]]) + hooks = MagicMock() + hooks.has_handlers.side_effect = lambda event: ( + event is HookEvent.USER_PROMPT_SUBMIT + ) + hooks.user_prompt_submit = AsyncMock( + return_value=UserPromptSubmitDecision( + event=HookEvent.USER_PROMPT_SUBMIT, + context=["replacement context"], + suppress_original_prompt=True, + ) + ) + hooks.take_session_context.return_value = () + state = TextualSessionState(thread_id="thread-1") + state.client_hooks = hooks + adapter = TextualUIAdapter( + mount_message=_mock_mount, + update_status=_noop_status, + request_approval=_mock_approval, + ) + + with patch( + "deepagents_code.tui.textual_adapter.dispatch_hook", + new_callable=AsyncMock, + ) as legacy: + await execute_task_textual( + user_input="secret prompt", + agent=agent, + assistant_id="assistant", + session_state=state, + adapter=adapter, + ) + + hooks.user_prompt_submit.assert_awaited_once() + stream_input = agent.stream_inputs[0] + assert isinstance(stream_input, dict) + assert stream_input["messages"] == [ + {"role": "system", "content": "replacement context"} + ] + assert not any( + call.args and call.args[0] in {"session.start", "user.prompt"} + for call in legacy.await_args_list + ) + + async def test_user_prompt_hook_block_prevents_stream(self) -> None: + from deepagents_code.app import TextualSessionState + from deepagents_code.hooks.client_lifecycle import ClientHookStopError + + agent = _SequencedAgent([[]]) + hooks = MagicMock() + hooks.has_handlers.side_effect = lambda event: ( + event is HookEvent.USER_PROMPT_SUBMIT + ) + hooks.user_prompt_submit = AsyncMock( + return_value=UserPromptSubmitDecision( + event=HookEvent.USER_PROMPT_SUBMIT, + continue_processing=False, + stop_reason="blocked prompt", + ) + ) + state = TextualSessionState(thread_id="thread-1") + state.client_hooks = hooks + adapter = TextualUIAdapter( + mount_message=_mock_mount, + update_status=_noop_status, + request_approval=_mock_approval, + ) + + with pytest.raises(ClientHookStopError, match="blocked prompt"): + await execute_task_textual( + user_input="blocked", + agent=agent, + assistant_id="assistant", + session_state=state, + adapter=adapter, + ) + + assert agent.stream_inputs == [] + + async def test_compact_session_start_precedes_continuation_with_model( + self, + ) -> None: + from deepagents_code.app import TextualSessionState + + order: list[str] = [] + agent = _SequencedAgent( + [ + [ + ( + (), + "messages", + ( + AIMessage(content="summary"), + {"lc_source": "summarization"}, + ), + ), + ((), "messages", (_text_message("continued"), {})), + ] + ] + ) + hooks = MagicMock() + hooks.has_handlers.return_value = False + hooks.take_session_context.return_value = () + hooks.pre_compact = AsyncMock() + hooks.session_start = AsyncMock( + side_effect=lambda *_args, **_kwargs: ( + order.append("start"), + SessionStartDecision(event=HookEvent.SESSION_START), + )[1] + ) + state = TextualSessionState(thread_id="thread-1") + state.client_hooks = hooks + + async def mount(widget: object) -> None: + await asyncio.sleep(0) + if isinstance(widget, AssistantMessage): + order.append("continuation") + + adapter = TextualUIAdapter( + mount_message=mount, + update_status=_noop_status, + request_approval=_mock_approval, + ) + with patch("deepagents_code.config.settings") as mock_settings: + mock_settings.model_name = "test:model" + mock_settings.model_provider = "test" + await execute_task_textual( + user_input="hello", + agent=agent, + assistant_id="assistant", + session_state=state, + adapter=adapter, + ) + + assert order[:2] == ["start", "continuation"] + hooks.pre_compact.assert_not_awaited() + await_args = hooks.session_start.await_args + assert await_args is not None + assert await_args.kwargs["model"] == "test:model" + + async def test_compact_permission_does_not_redispatch_precompact(self) -> None: + from deepagents_code.app import TextualSessionState + + order: list[str] = [] + request = {"name": "compact_conversation", "args": {}} + agent = _SequencedAgent( + [ + [ + _hitl_interrupt_chunk( + { + "action_requests": [request], + "review_configs": [], + } + ) + ], + [], + ] + ) + hooks = MagicMock() + hooks.has_handlers.side_effect = lambda event: ( + event + in { + HookEvent.PRE_COMPACT, + HookEvent.PERMISSION_REQUEST, + } + ) + hooks.take_session_context.return_value = () + hooks.pre_compact = AsyncMock() + hooks.permission_request = AsyncMock( + side_effect=lambda *_args: ( + order.append("permission"), + PermissionRequestDecision( + event=HookEvent.PERMISSION_REQUEST, + permission=PermissionEffect(behavior="allow"), + ), + )[1] + ) + state = TextualSessionState(thread_id="thread-1") + state.client_hooks = hooks + approval = AsyncMock() + adapter = TextualUIAdapter( + mount_message=_mock_mount, + update_status=_noop_status, + request_approval=approval, + ) + + await execute_task_textual( + user_input="compact", + agent=agent, + assistant_id="assistant", + session_state=state, + adapter=adapter, + ) + + assert order == ["permission"] + hooks.pre_compact.assert_not_awaited() + approval.assert_not_awaited() + + class TestExecuteTaskTextualAutoApproveInput: """Auto-approve must ride on run context, never a first-turn `Command`.""" From 6b3b52c4c2f7a9bb2f73969d700cc4ef0fd91c4a Mon Sep 17 00:00:00 2001 From: Johannes du Plessis Date: Fri, 24 Jul 2026 13:53:24 -0700 Subject: [PATCH 14/14] cr --- libs/code/deepagents_code/agent.py | 24 ++--- .../hooks/server_middleware.py | 90 ++++++++++++------- .../unit_tests/hooks/test_server_lifecycle.py | 52 ++++++++++- libs/code/tests/unit_tests/test_agent.py | 44 +++++++++ 4 files changed, 167 insertions(+), 43 deletions(-) diff --git a/libs/code/deepagents_code/agent.py b/libs/code/deepagents_code/agent.py index 7af98f54eb..af532e5717 100644 --- a/libs/code/deepagents_code/agent.py +++ b/libs/code/deepagents_code/agent.py @@ -2735,16 +2735,6 @@ def _subagent_cli_middleware( if restrictive_shell_allow_list is not None: agent_middleware.append(ShellAllowListMiddleware(restrictive_shell_allow_list)) - # Server-owned Hooks v2 lifecycle events (Pre/Post tool, Stop, subagent). - # Gated at runtime by `hooks_server_events` on the per-run context so idle - # sessions without configured handlers pay no interrupt round-trip. - from deepagents_code.hooks.server_middleware import ServerHooksMiddleware - - hooks_cwd = Path(effective_cwd) if effective_cwd is not None else Path.cwd() - if resolved_interrupt_on is not None: - agent_middleware.append(AsyncApprovalHITLMiddleware(resolved_interrupt_on)) - agent_middleware.append(ServerHooksMiddleware(cwd=hooks_cwd, mcp_tools=mcp_tools)) - # Get or use custom system prompt if system_prompt is None: system_prompt = get_system_prompt( @@ -2827,6 +2817,20 @@ def _subagent_cli_middleware( trusted_compaction_tool=compaction_middleware.tools[0], ) ) + elif resolved_interrupt_on is not None: + # `AutoModeHITLMiddleware` reports the same `HumanInTheLoopMiddleware` + # name, so installing both would trip `create_agent`'s duplicate-name + # assertion. Auto mode's specialized replacement wins when active. + agent_middleware.append(AsyncApprovalHITLMiddleware(resolved_interrupt_on)) + + # Server-owned Hooks v2 lifecycle events (Pre/Post tool, Stop, subagent). + # Gated at runtime by `hooks_server_events` on the per-run context so idle + # sessions without configured handlers pay no interrupt round-trip. Appended + # after the HITL middleware so `PreToolUse` resolves before approval routing. + from deepagents_code.hooks.server_middleware import ServerHooksMiddleware + + hooks_cwd = Path(effective_cwd) if effective_cwd is not None else Path.cwd() + agent_middleware.append(ServerHooksMiddleware(cwd=hooks_cwd, mcp_tools=mcp_tools)) if fs_tools is not None: # `fs_tools` is an explicit allowlist here (`--allow-fs-tools all` and an diff --git a/libs/code/deepagents_code/hooks/server_middleware.py b/libs/code/deepagents_code/hooks/server_middleware.py index 6c639ad5d4..c8771ddd47 100644 --- a/libs/code/deepagents_code/hooks/server_middleware.py +++ b/libs/code/deepagents_code/hooks/server_middleware.py @@ -13,7 +13,16 @@ from collections.abc import Mapping, Sequence from dataclasses import dataclass, field, replace from datetime import UTC, datetime, timedelta -from typing import TYPE_CHECKING, Any, Literal, NotRequired, TypeAlias, TypeVar, cast +from typing import ( + TYPE_CHECKING, + Any, + Literal, + NotRequired, + TypeAlias, + TypeGuard, + TypeVar, + cast, +) from uuid import UUID, uuid5 from langchain.agents.middleware.human_in_the_loop import ( @@ -183,7 +192,7 @@ def wrap_tool_call( request.runtime.context, request.runtime.config, self._cwd ) if pre.blocked is not None: - return _append_message_text(pre.blocked, pre.context) + return _append_message_text(pre.blocked, pre.context, call.id) started_or_blocked = self._maybe_subagent_start(request, call, context, gate) if isinstance(started_or_blocked, ToolMessage): return started_or_blocked @@ -191,7 +200,7 @@ def wrap_tool_call( started = time.perf_counter() result = handler(request) duration_ms = int((time.perf_counter() - started) * 1000) - result = _append_message_text(result, pre.context) + result = _append_message_text(result, pre.context, call.id) result = self._maybe_post_tool_use( call, context, gate, request.runtime.config, result, duration_ms ) @@ -216,7 +225,7 @@ async def awrap_tool_call( request.runtime.context, request.runtime.config, self._cwd ) if pre.blocked is not None: - return _append_message_text(pre.blocked, pre.context) + return _append_message_text(pre.blocked, pre.context, call.id) started_or_blocked = self._maybe_subagent_start(request, call, context, gate) if isinstance(started_or_blocked, ToolMessage): return started_or_blocked @@ -224,7 +233,7 @@ async def awrap_tool_call( started = time.perf_counter() result = await handler(request) duration_ms = int((time.perf_counter() - started) * 1000) - result = _append_message_text(result, pre.context) + result = _append_message_text(result, pre.context, call.id) result = self._maybe_post_tool_use( call, context, gate, request.runtime.config, result, duration_ms ) @@ -356,7 +365,7 @@ def _maybe_post_tool_use( ) -> ToolMessage | Command[Any]: if not _event_enabled(gate, HookEvent.POST_TOOL_USE): return result - if _tool_result_failed(result): + if _tool_result_failed(result, call.id): return result decision = _invoke_hook( context, @@ -371,7 +380,7 @@ def _maybe_post_tool_use( deadline=self._default_deadline, ) decision = _require_decision(decision, PostToolUseDecision) - return _apply_post_tool_use(result, decision) + return _apply_post_tool_use(result, decision, call.id) def _maybe_subagent_stop( self, @@ -392,14 +401,14 @@ def _maybe_subagent_stop( event=HookEvent.SUBAGENT_STOP, agent=agent, continuation_count=0, - last_assistant_message=_tool_result_text(result), + last_assistant_message=_tool_result_text(result, call.id), ), gate=gate, config=config, deadline=self._default_deadline, ) decision = _require_decision(decision, SubagentStopDecision) - return _apply_subagent_stop(result, decision) + return _apply_subagent_stop(result, decision, call.id) def _after_agent( self, @@ -796,15 +805,17 @@ def _ask_permission_via_hitl( def _append_message_text( result: ToolMessage | Command[Any], parts: Sequence[str], + call_id: str, ) -> ToolMessage | Command[Any]: if not parts: return result - return _append_tool_result_text(result, "\n".join(parts)) + return _append_tool_result_text(result, "\n".join(parts), call_id) def _apply_post_tool_use( result: ToolMessage | Command[Any], decision: PostToolUseDecision, + call_id: str, ) -> ToolMessage | Command[Any]: extras: list[str] = [] if decision.feedback: @@ -818,34 +829,34 @@ def _apply_post_tool_use( return _append_tool_result_text( result, "\n\n".join(part for part in extras if part), + call_id, ) def _apply_subagent_stop( result: ToolMessage | Command[Any], decision: SubagentStopDecision, + call_id: str, ) -> ToolMessage | Command[Any]: if not decision.context: return result - return _append_tool_result_text(result, "\n".join(decision.context)) + return _append_tool_result_text(result, "\n".join(decision.context), call_id) def _append_tool_result_text( result: ToolMessage | Command[Any], suffix: str, + call_id: str, ) -> ToolMessage | Command[Any]: if isinstance(result, ToolMessage): return _merge_tool_message_content(result, suffix) update = result.update if not isinstance(update, Mapping): return result - raw_messages = update.get("messages") - if not isinstance(raw_messages, Sequence) or isinstance(raw_messages, str): - return result changed = False messages: list[object] = [] - for message in raw_messages: - if isinstance(message, ToolMessage): + for message in _command_messages(result): + if _is_call_result(message, call_id): messages.append(_merge_tool_message_content(message, suffix)) changed = True else: @@ -855,19 +866,40 @@ def _append_tool_result_text( return replace(result, update={**update, "messages": messages}) -def _tool_result_failed(result: ToolMessage | Command[Any]) -> bool: +def _tool_result_failed(result: ToolMessage | Command[Any], call_id: str) -> bool: if isinstance(result, ToolMessage): return result.status == "error" + return any( + _is_call_result(message, call_id) and message.status == "error" + for message in _command_messages(result) + ) + + +def _command_messages(result: Command[Any]) -> Sequence[object]: + """Return the `messages` list carried by a `Command` update. + + Returns: + The update's messages, or an empty sequence when absent or malformed. + """ update = result.update if not isinstance(update, Mapping): - return False + return () messages = update.get("messages") if not isinstance(messages, Sequence) or isinstance(messages, str): - return False - return any( - isinstance(message, ToolMessage) and message.status == "error" - for message in messages - ) + return () + return messages + + +def _is_call_result(message: object, call_id: str) -> TypeGuard[ToolMessage]: + """Check whether a message is the `ToolMessage` for the in-flight call. + + A `Command` update may carry results for several calls, so hook context must + only read from and write to the one this wrapper is handling. + + Returns: + `True` when the message answers `call_id`. + """ + return isinstance(message, ToolMessage) and message.tool_call_id == call_id def _merge_tool_message_content(result: ToolMessage, suffix: str) -> ToolMessage: @@ -923,18 +955,14 @@ def _task_agent_identity(call: ToolCallData) -> AgentIdentity: return AgentIdentity(id=call.id or name, name=name) -def _tool_result_text(result: ToolMessage | Command[Any]) -> str: +def _tool_result_text(result: ToolMessage | Command[Any], call_id: str) -> str: if isinstance(result, ToolMessage): content = result.content return content if isinstance(content, str) else str(content) - update = result.update - if not isinstance(update, Mapping): - return "" - messages = update.get("messages") - if not isinstance(messages, Sequence) or isinstance(messages, str): - return "" return "\n".join( - str(message.content) for message in messages if isinstance(message, ToolMessage) + str(message.content) + for message in _command_messages(result) + if _is_call_result(message, call_id) ) diff --git a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py index f687cc5e63..5e31b6a20b 100644 --- a/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py +++ b/libs/code/tests/unit_tests/hooks/test_server_lifecycle.py @@ -7,7 +7,7 @@ import sys from datetime import UTC, datetime, timedelta from pathlib import Path -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any from unittest.mock import MagicMock from uuid import uuid4 @@ -54,6 +54,7 @@ ServerHooksMiddleware, ServerHooksState, _append_message_text, + _append_tool_result_text, _apply_post_tool_use, _apply_subagent_stop, _ask_permission_via_hitl, @@ -61,6 +62,8 @@ _invoke_hook, _merge_tool_message_content, _session_gate, + _tool_result_failed, + _tool_result_text, ) from deepagents_code.hooks.snapshot import HooksSnapshot @@ -254,6 +257,7 @@ def test_apply_subagent_stop_preserves_structured_content() -> None: event=HookEvent.SUBAGENT_STOP, context=["extra"], ), + "c1", ) assert isinstance(updated, ToolMessage) assert isinstance(updated.content, list) @@ -269,6 +273,7 @@ def test_apply_post_tool_use_appends_feedback_and_context() -> None: feedback=["fix it"], context=["note"], ), + "c1", ) assert "ok" in str(updated.content) assert "fix it" in str(updated.content) @@ -322,6 +327,49 @@ def test_post_tool_use_updates_successful_command_result( invoke.assert_called_once() +def _multi_result_command() -> Command[Any]: + return Command( + update={ + "messages": [ + ToolMessage(content="mine", name="execute", tool_call_id="c1"), + ToolMessage( + content="theirs", + name="execute", + tool_call_id="c2", + status="error", + ), + ] + } + ) + + +def test_append_tool_result_text_only_touches_matching_call() -> None: + updated = _append_tool_result_text(_multi_result_command(), "hook context", "c1") + + assert isinstance(updated, Command) + assert isinstance(updated.update, dict) + mine, theirs = updated.update["messages"] + assert "hook context" in str(mine.content) + assert str(theirs.content) == "theirs" + + +def test_append_tool_result_text_leaves_command_without_matching_call() -> None: + result = _multi_result_command() + + assert _append_tool_result_text(result, "hook context", "c3") is result + + +def test_tool_result_text_reads_only_matching_call() -> None: + assert _tool_result_text(_multi_result_command(), "c1") == "mine" + + +def test_tool_result_failed_ignores_unrelated_failure() -> None: + result = _multi_result_command() + + assert _tool_result_failed(result, "c1") is False + assert _tool_result_failed(result, "c2") is True + + def test_post_tool_use_skips_failed_tool_message( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -357,7 +405,7 @@ def test_post_tool_use_skips_failed_tool_message( def test_append_pretool_context_to_result() -> None: result = ToolMessage(content="ran", tool_call_id="c1", name="execute") - updated = _append_message_text(result, ("pre context",)) + updated = _append_message_text(result, ("pre context",), "c1") assert isinstance(updated, ToolMessage) assert "ran" in str(updated.content) assert "pre context" in str(updated.content) diff --git a/libs/code/tests/unit_tests/test_agent.py b/libs/code/tests/unit_tests/test_agent.py index 149953a2cb..90161ce767 100644 --- a/libs/code/tests/unit_tests/test_agent.py +++ b/libs/code/tests/unit_tests/test_agent.py @@ -4797,6 +4797,50 @@ def test_auto_mode_enabled_wires_middleware(self, tmp_path: Path) -> None: compaction_middleware ) + @pytest.mark.parametrize("auto_mode_enabled", [True, False]) + def test_single_hitl_slot_precedes_server_hooks( + self, + tmp_path: Path, + *, + auto_mode_enabled: bool, + ) -> None: + """One HITL middleware is installed, ahead of the server hook middleware. + + `AutoModeHITLMiddleware` reports the stock `HumanInTheLoopMiddleware` + name, so pairing it with the standalone approval middleware would trip + `create_agent`'s duplicate-name assertion. `ServerHooksMiddleware` must + stay behind whichever one is installed so its `after_model` `PreToolUse` + pass resolves before approval routing. + """ + from deepagents_code.hooks.server_middleware import ServerHooksMiddleware + + middleware = self._capture_middleware( + tmp_path, auto_mode_enabled=auto_mode_enabled + ) + + hitl = [item for item in middleware if item.name == "HumanInTheLoopMiddleware"] + hooks = next( + item for item in middleware if isinstance(item, ServerHooksMiddleware) + ) + + assert len(hitl) == 1 + assert middleware.index(hitl[0]) < middleware.index(hooks) + + def test_auto_mode_agent_builds(self, tmp_path: Path) -> None: + """Auto mode compiles a real graph rather than aborting on duplicates.""" + agent, _backend = create_cli_agent( + model=_make_fake_chat_model(), + assistant_id="test-agent", + enable_memory=False, + enable_skills=False, + enable_shell=False, + system_prompt="test prompt", + cwd=tmp_path, + auto_mode_enabled=True, + ) + + assert agent is not None + def test_auto_mode_omitted_outside_interactive(self, tmp_path: Path) -> None: """Auto is refused (no middleware) in a non-interactive session.""" from deepagents_code.auto_mode import AutoModeHITLMiddleware