refactor(acp): events _upgrade_queue single key; edit_approval uses Path.is_relative_to; session toolset expansion via dict.fromkeys; fork_session early return

This commit is contained in:
Teknium
2026-09-02 18:13:29 -07:00
parent 425a9d10b6
commit df7c867027
4 changed files with 24 additions and 39 deletions
+4 -12
View File
@@ -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
+8 -7
View File
@@ -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, {})
+7 -9
View File
@@ -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(
+5 -11
View File
@@ -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 = [