diff --git a/backend/apps/agents/agent_manager.py b/backend/apps/agents/agent_manager.py index 1922a773..e8e9554c 100644 --- a/backend/apps/agents/agent_manager.py +++ b/backend/apps/agents/agent_manager.py @@ -59,6 +59,7 @@ from backend.apps.agents.manager import context_budget from backend.apps.agents.manager.streaming.state import ThinkingState, TurnState from backend.apps.agents.manager.streaming.hook_context import HookContext from backend.apps.agents.manager.streaming import thinking as thinking_mod +from backend.apps.agents.manager.streaming import tool_result_hook from backend.apps.agents.manager.permissions import gate_hooks from backend.apps.agents.manager.view_builder_state import ( VIEW_BUILDER_RENDER_MAX_RETRIES, @@ -79,7 +80,6 @@ from backend.apps.agents.manager.session.history_compaction import ( _build_history_prefix, _estimate_post_compact_input, _get_branch_messages, - _truncate_large_tool_result, ) from backend.apps.agents.manager.prompt.prompt_context import ( _build_browser_context, @@ -417,6 +417,7 @@ class AgentManager: prompt=prompt, builtin_perms=builtin_perms, policy_defaults={}, + sessions=self.sessions, ) async def can_use_tool(tool_name, input_data, context): @@ -426,221 +427,7 @@ class AgentManager: return await gate_hooks.pre_tool_hook(hook_ctx, input_data, tool_use_id, context) async def post_tool_hook(input_data, tool_use_id, context): - elapsed_ms = None - if tool_use_id and tool_use_id in hook_ctx.tool_start_times: - elapsed_ms = int((time.time() - hook_ctx.tool_start_times.pop(tool_use_id)) * 1000) - - raw_response = input_data.get("tool_response", "") - - # Track individual tool execution - hook_tool_name_early = input_data.get("tool_name", "") - if hook_tool_name_early: - _is_mcp = "__" in hook_tool_name_early - _mcp_server = "" - _tool_short = hook_tool_name_early - if _is_mcp: - _mcp_match = re.match(r"mcp__([^_]+(?:-[^_]+)*)__(.+)", hook_tool_name_early) - if _mcp_match: - _mcp_server = _mcp_match.group(1) - _tool_short = _mcp_match.group(2) - - # Accumulate per-tool latency on the session. Lets the - # cloud aggregate a tool-latency distribution into the - # existing daily.summary without firing per-tool events. - if elapsed_ms is not None and elapsed_ms >= 0: - latencies = getattr(session, "tool_latencies", None) - if latencies is None: - latencies = {} - try: - session.tool_latencies = latencies - except Exception: - latencies = None - if latencies is not None: - slot = latencies.get(hook_tool_name_early) - if slot is None: - slot = {"count": 0, "total_ms": 0, "max_ms": 0} - latencies[hook_tool_name_early] = slot - slot["count"] = slot.get("count", 0) + 1 - slot["total_ms"] = slot.get("total_ms", 0) + elapsed_ms - slot["max_ms"] = max(slot.get("max_ms", 0), elapsed_ms) - - # Determine tool success - _tool_success = True - if isinstance(raw_response, str): - _tool_success = not (raw_response.startswith("Error") or raw_response.startswith("Traceback")) - elif isinstance(raw_response, dict): - _tool_success = "error" not in raw_response - elif isinstance(raw_response, list): - _tool_success = len(raw_response) > 0 - - - if isinstance(raw_response, list) and raw_response: - text_parts = [ - block.get("text", "") - for block in raw_response - if isinstance(block, dict) and block.get("type") == "text" - ] - if text_parts: - raw_response = "\n".join(text_parts) if len(text_parts) > 1 else text_parts[0] - - if isinstance(raw_response, str): - content = raw_response - else: - try: - import json as _json - content = _json.dumps(raw_response, indent=2, default=str) - except Exception: - content = str(raw_response) - - hook_tool_name_for_errors = input_data.get("tool_name", "") - wrote_files = hook_tool_name_for_errors in ("Write", "Edit", "MultiEdit") - tool_in = input_data.get("tool_input") or {} - file_path = tool_in.get("file_path") or tool_in.get("path") or "" - wrote_frontend_file = wrote_files and "/frontend/" in file_path - installed_pkg = False - if hook_tool_name_for_errors == "Bash": - bash_in = input_data.get("tool_input") or {} - cmd = (bash_in.get("command") or "").lower() - installed_pkg = any(s in cmd for s in ( - "npm install", "npm i ", "npm uninstall", "npm ci", - "pnpm add", "pnpm install", "pnpm remove", - "yarn add", "yarn install", "yarn remove", - )) - - if session.mode == "view-builder" and (wrote_frontend_file or installed_pkg): - view_builder_dirty_sessions.add(session.id) - try: - from backend.apps.outputs.runtime import ( - manager as outputs_runtime_manager, - ) - outputs_runtime_manager.reset_render_state_for_workspace(session.id) - except Exception: - pass - elif wrote_files: - if file_path: - try: - await asyncio.sleep(0.4) - from backend.apps.outputs.runtime import ( - manager as outputs_runtime_manager, - ) - errs = outputs_runtime_manager.drain_errors_for_path(file_path) - except Exception: - errs = [] - if errs: - joined = "\n".join(errs[-20:]) - content = ( - f"{content}\n\n" - f"---\nBuild server reported (after this write):\n{joined}" - ) - - result_payload = {"text": content} - hook_tool_name = input_data.get("tool_name", "") - if hook_tool_name: - result_payload["tool_name"] = hook_tool_name - if elapsed_ms is not None: - result_payload["elapsed_ms"] = elapsed_ms - - if hook_tool_name == "Agent": - tool_input = input_data.get("tool_input", {}) - agent_prompt = tool_input.get("prompt", tool_input.get("task", "")) - - sub_text = content - sub_cost = 0.0 - sub_tokens = {"input": 0, "output": 0} - sub_model = session.model - if isinstance(raw_response, dict): - blocks = raw_response.get("content") - if isinstance(blocks, list): - parts = [ - b.get("text", "") - for b in blocks - if isinstance(b, dict) and b.get("type") == "text" - ] - if parts: - sub_text = "\n".join(parts) if len(parts) > 1 else parts[0] - elif isinstance(raw_response.get("text"), str): - sub_text = raw_response["text"] - usage = raw_response.get("usage", {}) - if isinstance(usage, dict): - sub_tokens["input"] = usage.get("input_tokens", 0) + usage.get("cache_creation_input_tokens", 0) + usage.get("cache_read_input_tokens", 0) - # Pill-only lane: NEW (uncached) input, excludes the cached - # static prefix so the bubble shows what this turn added. - sub_tokens["input_fresh"] = usage.get("input_tokens", 0) - sub_tokens["output"] = usage.get("output_tokens", 0) - if raw_response.get("total_cost_usd"): - sub_cost = raw_response["total_cost_usd"] - if raw_response.get("model"): - sub_model = raw_response["model"] - - sub_session_id = uuid4().hex - sub_name = agent_prompt[:50] if agent_prompt else "Sub-agent" - # Subagent context isolation invariant (Phase 3, Layer P): - # children DO NOT inherit the parent's active_mcps or - # compaction state. They start with the AgentSession - # defaults (empty lists). Reasoning: - # - Security: a parent that activated Gmail shouldn't - # leak Gmail tools to a subagent doing an unrelated - # task. The user only approved Gmail for the parent. - # - Token cost: subagents typically have a narrow task, - # they don't need the parent's full activated set. - # - Failure isolation: if the parent compacted history, - # the subagent shouldn't inherit a summary it can't - # re-expand. - # If a subagent ever needs a parent activation, the user - # must approve it explicitly via MCPActivate inside the - # subagent session, same gate as a fresh top-level chat. - sub_session = AgentSession( - id=sub_session_id, - name=sub_name, - status="completed", - model=sub_model, - mode="sub-agent", - cwd=session.cwd, - created_at=datetime.now(), - cost_usd=sub_cost, - tokens=sub_tokens, - messages=[ - Message(role="user", content=agent_prompt, branch_id="main"), - Message(role="assistant", content=sub_text, branch_id="main"), - ], - dashboard_id=session.dashboard_id, - parent_session_id=session_id, - # Explicit empty list (matches the model default) so - # the invariant is visible at the spawn site rather - # than relying on the field's default_factory. - active_mcps=[], - ) - apply_context_window(sub_session) - self.sessions[sub_session_id] = sub_session - await ws_manager.broadcast_global("agent:status", { - "session_id": sub_session_id, - "status": sub_session.status, - "session": sub_session.model_dump(mode="json"), - }) - result_payload["sub_session_id"] = sub_session_id - - result_msg = Message(role="tool_result", content=result_payload, branch_id=session.active_branch_id) - # Spill oversized tool results to per-session disk storage. - # The replacement keeps the first 4KB inline so the model - # retains some signal; the rest lives on disk for the UI to - # surface in the compaction drawer. Crucially this happens - # at *write* time (before the next turn ships history to the - # SDK) so the bloat never re-enters context. - try: - truncated_content, blob_path = _truncate_large_tool_result( - result_msg.content, session.id, result_msg.id - ) - if blob_path: - result_msg.content = truncated_content - logger.info(f"Spilled tool result {result_msg.id} ({len(blob_path)} chars) to {blob_path}") - except Exception: - logger.exception("Tool result truncation failed; keeping inline body") - session.messages.append(result_msg) - await ws_manager.send_to_session(session_id, "agent:message", { - "session_id": session_id, - "message": result_msg.model_dump(mode="json"), - }) - return {"continue_": True} + return await tool_result_hook.post_tool_hook(hook_ctx, input_data, tool_use_id, context) try: _, mode_sys_prompt, _ = self._resolve_mode(session.mode) diff --git a/backend/apps/agents/manager/streaming/hook_context.py b/backend/apps/agents/manager/streaming/hook_context.py index e433d886..acc512c1 100644 --- a/backend/apps/agents/manager/streaming/hook_context.py +++ b/backend/apps/agents/manager/streaming/hook_context.py @@ -6,7 +6,7 @@ pending_approvals are visible to the loop.""" from typing import Dict -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, InstanceOf from backend.apps.agents.core.models import AgentSession @@ -19,6 +19,9 @@ class HookContext(BaseModel): prompt: str builtin_perms: Dict[str, str] policy_defaults: Dict[str, str] + # The manager's LIVE session registry (InstanceOf keeps the reference, so a sub-agent + # the post hook spawns is visible to the manager; a plain Dict field pydantic would copy). + sessions: InstanceOf[dict] # tool_use_id -> wall-clock start (s); pre records it, post pops it for elapsed_ms. tool_start_times: Dict[str, float] = {} # Consecutive ToolSearch calls; a run of these is the "looping on ToolSearch" wedge. diff --git a/backend/apps/agents/manager/streaming/tool_result_hook.py b/backend/apps/agents/manager/streaming/tool_result_hook.py new file mode 100644 index 00000000..5e1f6439 --- /dev/null +++ b/backend/apps/agents/manager/streaming/tool_result_hook.py @@ -0,0 +1,225 @@ +"""The SDK PostToolUse hook, lifted out of the agent loop. Runs after every tool call: +records per-tool latency, normalizes the raw tool response into displayable text, re-renders +view-builder writes (and drains build errors), materializes a spawned Agent sub-session into +the manager registry, spills oversized results to disk, and broadcasts the tool_result message. +Operates on the HookContext (its `sessions` is the manager's live registry). The dict returns +and payloads are the SDK hook protocol / existing message shapes, not internal models.""" + +import asyncio +import logging +import time +from datetime import datetime +from typing import Dict +from uuid import uuid4 + +from typeguard import typechecked + +from backend.apps.agents.core.models import AgentSession, Message +from backend.apps.agents.core.ws_manager import ws_manager +from backend.apps.agents.manager.session.apply_context_window import apply_context_window +from backend.apps.agents.manager.session.history_compaction import _truncate_large_tool_result +from backend.apps.agents.manager.streaming.hook_context import HookContext +from backend.apps.agents.manager.view_builder_state import view_builder_dirty_sessions + +logger = logging.getLogger(__name__) + + +@typechecked +async def post_tool_hook(ctx: HookContext, input_data: dict, tool_use_id, context) -> Dict[str, object]: + session = ctx.session + session_id = ctx.session_id + + elapsed_ms = None + if tool_use_id and tool_use_id in ctx.tool_start_times: + elapsed_ms = int((time.time() - ctx.tool_start_times.pop(tool_use_id)) * 1000) + + raw_response = input_data.get("tool_response", "") + + # Accumulate per-tool latency on the session. Lets the cloud aggregate a + # tool-latency distribution into the existing daily.summary without firing + # per-tool events. + hook_tool_name_early = input_data.get("tool_name", "") + if hook_tool_name_early and elapsed_ms is not None and elapsed_ms >= 0: + latencies = getattr(session, "tool_latencies", None) + if latencies is None: + latencies = {} + try: + session.tool_latencies = latencies + except Exception: + latencies = None + if latencies is not None: + slot = latencies.get(hook_tool_name_early) + if slot is None: + slot = {"count": 0, "total_ms": 0, "max_ms": 0} + latencies[hook_tool_name_early] = slot + slot["count"] = slot.get("count", 0) + 1 + slot["total_ms"] = slot.get("total_ms", 0) + elapsed_ms + slot["max_ms"] = max(slot.get("max_ms", 0), elapsed_ms) + + if isinstance(raw_response, list) and raw_response: + text_parts = [ + block.get("text", "") + for block in raw_response + if isinstance(block, dict) and block.get("type") == "text" + ] + if text_parts: + raw_response = "\n".join(text_parts) if len(text_parts) > 1 else text_parts[0] + + if isinstance(raw_response, str): + content = raw_response + else: + try: + import json as json_lib + content = json_lib.dumps(raw_response, indent=2, default=str) + except Exception: + content = str(raw_response) + + hook_tool_name_for_errors = input_data.get("tool_name", "") + wrote_files = hook_tool_name_for_errors in ("Write", "Edit", "MultiEdit") + tool_in = input_data.get("tool_input") or {} + file_path = tool_in.get("file_path") or tool_in.get("path") or "" + wrote_frontend_file = wrote_files and "/frontend/" in file_path + installed_pkg = False + if hook_tool_name_for_errors == "Bash": + bash_in = input_data.get("tool_input") or {} + cmd = (bash_in.get("command") or "").lower() + installed_pkg = any(s in cmd for s in ( + "npm install", "npm i ", "npm uninstall", "npm ci", + "pnpm add", "pnpm install", "pnpm remove", + "yarn add", "yarn install", "yarn remove", + )) + + if session.mode == "view-builder" and (wrote_frontend_file or installed_pkg): + view_builder_dirty_sessions.add(session.id) + try: + from backend.apps.outputs.runtime import ( + manager as outputs_runtime_manager, + ) + outputs_runtime_manager.reset_render_state_for_workspace(session.id) + except Exception: + pass + elif wrote_files: + if file_path: + try: + await asyncio.sleep(0.4) + from backend.apps.outputs.runtime import ( + manager as outputs_runtime_manager, + ) + errs = outputs_runtime_manager.drain_errors_for_path(file_path) + except Exception: + errs = [] + if errs: + joined = "\n".join(errs[-20:]) + content = ( + f"{content}\n\n" + f"---\nBuild server reported (after this write):\n{joined}" + ) + + result_payload = {"text": content} + hook_tool_name = input_data.get("tool_name", "") + if hook_tool_name: + result_payload["tool_name"] = hook_tool_name + if elapsed_ms is not None: + result_payload["elapsed_ms"] = elapsed_ms + + if hook_tool_name == "Agent": + tool_input = input_data.get("tool_input", {}) + agent_prompt = tool_input.get("prompt", tool_input.get("task", "")) + + sub_text = content + sub_cost = 0.0 + sub_tokens = {"input": 0, "output": 0} + sub_model = session.model + if isinstance(raw_response, dict): + blocks = raw_response.get("content") + if isinstance(blocks, list): + parts = [ + b.get("text", "") + for b in blocks + if isinstance(b, dict) and b.get("type") == "text" + ] + if parts: + sub_text = "\n".join(parts) if len(parts) > 1 else parts[0] + elif isinstance(raw_response.get("text"), str): + sub_text = raw_response["text"] + usage = raw_response.get("usage", {}) + if isinstance(usage, dict): + sub_tokens["input"] = usage.get("input_tokens", 0) + usage.get("cache_creation_input_tokens", 0) + usage.get("cache_read_input_tokens", 0) + # Pill-only lane: NEW (uncached) input, excludes the cached + # static prefix so the bubble shows what this turn added. + sub_tokens["input_fresh"] = usage.get("input_tokens", 0) + sub_tokens["output"] = usage.get("output_tokens", 0) + if raw_response.get("total_cost_usd"): + sub_cost = raw_response["total_cost_usd"] + if raw_response.get("model"): + sub_model = raw_response["model"] + + sub_session_id = uuid4().hex + sub_name = agent_prompt[:50] if agent_prompt else "Sub-agent" + # Subagent context isolation invariant (Phase 3, Layer P): + # children DO NOT inherit the parent's active_mcps or + # compaction state. They start with the AgentSession + # defaults (empty lists). Reasoning: + # - Security: a parent that activated Gmail shouldn't + # leak Gmail tools to a subagent doing an unrelated + # task. The user only approved Gmail for the parent. + # - Token cost: subagents typically have a narrow task, + # they don't need the parent's full activated set. + # - Failure isolation: if the parent compacted history, + # the subagent shouldn't inherit a summary it can't + # re-expand. + # If a subagent ever needs a parent activation, the user + # must approve it explicitly via MCPActivate inside the + # subagent session, same gate as a fresh top-level chat. + sub_session = AgentSession( + id=sub_session_id, + name=sub_name, + status="completed", + model=sub_model, + mode="sub-agent", + cwd=session.cwd, + created_at=datetime.now(), + cost_usd=sub_cost, + tokens=sub_tokens, + messages=[ + Message(role="user", content=agent_prompt, branch_id="main"), + Message(role="assistant", content=sub_text, branch_id="main"), + ], + dashboard_id=session.dashboard_id, + parent_session_id=session_id, + # Explicit empty list (matches the model default) so + # the invariant is visible at the spawn site rather + # than relying on the field's default_factory. + active_mcps=[], + ) + apply_context_window(sub_session) + ctx.sessions[sub_session_id] = sub_session + await ws_manager.broadcast_global("agent:status", { + "session_id": sub_session_id, + "status": sub_session.status, + "session": sub_session.model_dump(mode="json"), + }) + result_payload["sub_session_id"] = sub_session_id + + result_msg = Message(role="tool_result", content=result_payload, branch_id=session.active_branch_id) + # Spill oversized tool results to per-session disk storage. + # The replacement keeps the first 4KB inline so the model + # retains some signal; the rest lives on disk for the UI to + # surface in the compaction drawer. Crucially this happens + # at *write* time (before the next turn ships history to the + # SDK) so the bloat never re-enters context. + try: + truncated_content, blob_path = _truncate_large_tool_result( + result_msg.content, session.id, result_msg.id + ) + if blob_path: + result_msg.content = truncated_content + logger.info(f"Spilled tool result {result_msg.id} ({len(blob_path)} chars) to {blob_path}") + except Exception: + logger.exception("Tool result truncation failed; keeping inline body") + session.messages.append(result_msg) + await ws_manager.send_to_session(session_id, "agent:message", { + "session_id": session_id, + "message": result_msg.model_dump(mode="json"), + }) + return {"continue_": True} diff --git a/backend/tests/test_gate_hooks.py b/backend/tests/test_gate_hooks.py index 8bc0c4fa..56749d9a 100644 --- a/backend/tests/test_gate_hooks.py +++ b/backend/tests/test_gate_hooks.py @@ -21,6 +21,7 @@ def _ctx() -> HookContext: prompt="hi", builtin_perms={}, policy_defaults={}, + sessions={}, ) diff --git a/backend/tests/test_tool_result_hook.py b/backend/tests/test_tool_result_hook.py new file mode 100644 index 00000000..1c0e3095 --- /dev/null +++ b/backend/tests/test_tool_result_hook.py @@ -0,0 +1,67 @@ +"""Unit coverage for the extracted PostToolUse hook (tool_result_hook). The streaming harness +mocks claude_agent_sdk.query, so it never fires the SDK's PostToolUse hooks; this pins the +behavior directly: a tool result becomes a tool_result message, and an Agent tool spawns a +sub-session into the manager's LIVE registry (the InstanceOf[dict] sharing, the subtle bit).""" + +import pytest +from unittest.mock import patch, AsyncMock + +from backend.apps.agents.core.models import AgentSession +from backend.apps.agents.manager.streaming.hook_context import HookContext +from backend.apps.agents.manager.streaming import tool_result_hook + + +def _ctx(registry: dict) -> HookContext: + session = AgentSession(name="t", model="sonnet", dashboard_id="d") + registry[session.id] = session + return HookContext( + session=session, + session_id=session.id, + prompt="hi", + builtin_perms={}, + policy_defaults={}, + sessions=registry, + ) + + +@pytest.mark.asyncio +async def test_normal_tool_result_appends_message_and_continues(): + registry: dict = {} + ctx = _ctx(registry) + before = len(ctx.session.messages) + with patch.object(tool_result_hook.ws_manager, "send_to_session", new=AsyncMock()) as send: + out = await tool_result_hook.post_tool_hook( + ctx, {"tool_name": "Read", "tool_response": "file body", "tool_input": {"file_path": "/x"}}, "tu1", None + ) + assert out == {"continue_": True} + assert len(ctx.session.messages) == before + 1 + msg = ctx.session.messages[-1] + assert msg.role == "tool_result" + assert "file body" in str(msg.content) + send.assert_awaited() # the tool_result is broadcast to the UI + + +@pytest.mark.asyncio +async def test_agent_tool_spawns_subsession_into_live_registry(): + registry: dict = {} + ctx = _ctx(registry) + parent_id = ctx.session_id + raw = { + "content": [{"type": "text", "text": "sub-agent did the work"}], + "usage": {"input_tokens": 7, "output_tokens": 3}, + "total_cost_usd": 0.01, + "model": "sonnet", + } + with patch.object(tool_result_hook.ws_manager, "send_to_session", new=AsyncMock()), \ + patch.object(tool_result_hook.ws_manager, "broadcast_global", new=AsyncMock()): + out = await tool_result_hook.post_tool_hook( + ctx, {"tool_name": "Agent", "tool_response": raw, "tool_input": {"prompt": "do x"}}, "tu1", None + ) + assert out == {"continue_": True} + # exactly one NEW session registered (besides the parent), parented correctly + children = [s for sid, s in registry.items() if sid != parent_id] + assert len(children) == 1 + child = children[0] + assert child.parent_session_id == parent_id + assert child.active_mcps == [] # context-isolation invariant: no inherited activations + assert "sub-agent did the work" in str(child.messages[-1].content)