From df7c867027ba9cf7ba8ed2b6a7814ffb377d1997 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 18:13:29 -0700 Subject: [PATCH] refactor(acp): events _upgrade_queue single key; edit_approval uses Path.is_relative_to; session toolset expansion via dict.fromkeys; fork_session early return --- acp_adapter/edit_approval.py | 16 ++++------------ acp_adapter/events.py | 15 ++++++++------- acp_adapter/server.py | 16 +++++++--------- acp_adapter/session.py | 16 +++++----------- 4 files changed, 24 insertions(+), 39 deletions(-) diff --git a/acp_adapter/edit_approval.py b/acp_adapter/edit_approval.py index a126140c9e..31f3889588 100644 --- a/acp_adapter/edit_approval.py +++ b/acp_adapter/edit_approval.py @@ -115,8 +115,8 @@ def _proposal_for_patch_v4a(arguments: dict[str, Any]) -> EditProposal: if not paths: raise ValueError("no file paths found in V4A patch") single = len(paths) == 1 - # ACP only supports a single diff payload: surface the exact V4A patch as - # new_text so patch-mode calls are permissioned and denied patches cannot mutate. + # ACP only supports a single diff payload: surface the exact V4A patch as new_text so + # patch-mode calls are permissioned and denied patches cannot mutate. return EditProposal( "patch", paths[0] if single else ", ".join(paths), _read_text_if_exists(paths[0]) if single else None, patch_body, dict(arguments), @@ -143,14 +143,6 @@ def _is_sensitive_auto_approve_path(path: str) -> bool: return bool(lowered & {".git", ".ssh"}) or Path(path).name.lower() in SENSITIVE_AUTO_APPROVE_NAMES -def _is_under(path: Path, root: Path) -> bool: - try: - path.relative_to(root) - return True - except ValueError: - return False - - def should_auto_approve_edit(proposal: EditProposal, policy: str, cwd: str | None = None) -> bool: """Return whether an ACP edit proposal may bypass the prompt for this session. @@ -165,10 +157,10 @@ def should_auto_approve_edit(proposal: EditProposal, policy: str, cwd: str | Non if policy == AUTO_APPROVE_WORKSPACE: # tempfile.gettempdir() is the real temp root on every platform # (``/private/tmp`` on macOS since resolve() follows the symlink). - if _is_under(path, Path(tempfile.gettempdir()).resolve(strict=False)): + if path.is_relative_to(Path(tempfile.gettempdir()).resolve(strict=False)): return True if cwd: - return _is_under(path, Path(cwd).expanduser().resolve(strict=False)) + return path.is_relative_to(Path(cwd).expanduser().resolve(strict=False)) return False diff --git a/acp_adapter/events.py b/acp_adapter/events.py index e9a89fc4e1..6c80c0cc23 100644 --- a/acp_adapter/events.py +++ b/acp_adapter/events.py @@ -65,12 +65,11 @@ def _send_update(conn: acp.Client, session_id: str, loop: asyncio.AbstractEventL logger.debug("Failed to send ACP update", exc_info=True) -def _upgrade_queue(tool_call_ids: Dict[str, Deque[str]], key: Any, store_key: Any) -> Deque[str] | None: +def _upgrade_queue(tool_call_ids: Dict[str, Deque[str]], name: str) -> Deque[str] | None: """Fetch the per-tool FIFO of pending call IDs, upgrading a legacy bare-string entry in place.""" - queue = tool_call_ids.get(key) + queue = tool_call_ids.get(name) if isinstance(queue, str): - queue = deque([queue]) - tool_call_ids[store_key] = queue + queue = tool_call_ids[name] = deque([queue]) return queue @@ -102,7 +101,7 @@ def make_tool_progress_cb( args = {} tc_id = make_tool_call_id() - queue = _upgrade_queue(tool_call_ids, name, name) + queue = _upgrade_queue(tool_call_ids, name) if queue is None: queue = tool_call_ids[name] = deque() queue.append(tc_id) @@ -174,8 +173,10 @@ def make_step_cb( elif isinstance(tool_info, str): tool_name = tool_info - queue = _upgrade_queue(tool_call_ids, tool_name or "", tool_name) - if not tool_name or not queue: + if not tool_name: + continue + queue = _upgrade_queue(tool_call_ids, tool_name) + if not queue: continue tc_id = queue.popleft() meta = tool_call_meta.pop(tc_id, {}) diff --git a/acp_adapter/server.py b/acp_adapter/server.py index a86dbbf7c9..da0c7c8021 100644 --- a/acp_adapter/server.py +++ b/acp_adapter/server.py @@ -1299,16 +1299,14 @@ class HermesACPAgent(acp.Agent): self, cwd: str, session_id: str, mcp_servers: list | None = None, **kwargs: Any ) -> ForkSessionResponse: state = self.session_manager.fork_session(session_id, cwd=cwd) - new_id = state.session_id if state else "" - if state is not None: - await self._register_session_mcp_servers(state, mcp_servers) - logger.info("Forked session %s -> %s", session_id, new_id) - if new_id: - self._schedule_available_commands_update(new_id) + if state is None: + logger.info("Forked session %s -> %s", session_id, "") + return ForkSessionResponse(session_id="") + await self._register_session_mcp_servers(state, mcp_servers) + logger.info("Forked session %s -> %s", session_id, state.session_id) + self._schedule_available_commands_update(state.session_id) return ForkSessionResponse( - session_id=new_id, - models=self._build_model_state(state) if state is not None else None, - modes=self._session_modes(state) if state is not None else None, + session_id=state.session_id, models=self._build_model_state(state), modes=self._session_modes(state) ) async def list_sessions( diff --git a/acp_adapter/session.py b/acp_adapter/session.py index 1871752ab6..fe01fb936c 100644 --- a/acp_adapter/session.py +++ b/acp_adapter/session.py @@ -104,15 +104,9 @@ def _register_task_cwd(task_id: str, cwd: str) -> None: def _expand_acp_enabled_toolsets(toolsets: List[str] | None = None, mcp_server_names: List[str] | None = None) -> List[str]: """Return ACP toolsets plus explicit MCP server toolsets for this session.""" - expanded: List[str] = [] - for name in list(toolsets or ["hermes-acp"]): - if name and name not in expanded: - expanded.append(name) - for server_name in list(mcp_server_names or []): - toolset_name = f"mcp-{server_name}" - if server_name and toolset_name not in expanded: - expanded.append(toolset_name) - return expanded + names = [n for n in (toolsets or ["hermes-acp"]) if n] + names += [f"mcp-{s}" for s in (mcp_server_names or []) if s] + return list(dict.fromkeys(names)) def _parse_model_config(mc: Any) -> dict: @@ -223,7 +217,7 @@ class SessionManager: seen_ids = set(self._sessions.keys()) results = [] for s in self._sessions.values(): - if len(s.history) <= 0 or not _matches(s.cwd): + if not s.history or not _matches(s.cwd): continue persisted = persisted_rows.get(s.session_id, {}) results.append(_session_info( @@ -387,7 +381,7 @@ class SessionManager: default_model, config_provider = "", None if isinstance(model_cfg, dict): default_model, config_provider = str(model_cfg.get("default") or ""), model_cfg.get("provider") - elif isinstance(model_cfg, str) and model_cfg.strip(): + elif isinstance(model_cfg, str): default_model = model_cfg.strip() configured_mcp_servers = [