refactor(acp): unify tool-arg coercion (coerce_tool_args) across events/server; inline history chunk factory; flatten cancel/reset/queue paths; entry flag dispatch

This commit is contained in:
Teknium
2026-09-02 18:29:26 -07:00
parent cbd2011900
commit 41e861bf71
4 changed files with 56 additions and 83 deletions

View File

@@ -80,10 +80,9 @@ def _load_env() -> None:
hermes_home = get_hermes_home()
loaded = load_hermes_dotenv(hermes_home=hermes_home)
log = logging.getLogger(__name__)
if loaded:
for env_file in loaded:
log.info("Loaded env from %s", env_file)
else:
for env_file in loaded or ():
log.info("Loaded env from %s", env_file)
if not loaded:
log.info("No .env found at %s, using system env", hermes_home / ".env")
@@ -164,12 +163,9 @@ def _run_setup_browser(assume_yes: bool = False) -> int:
def main(argv: list[str] | None = None) -> None:
"""Entry point: load env, configure logging, run the ACP agent."""
args = _parse_args(argv)
if args.version:
return _print_version()
if args.check:
return _run_check()
if args.setup:
return _run_setup()
for flag, action in (("version", _print_version), ("check", _run_check), ("setup", _run_setup)):
if getattr(args, flag):
return action()
if args.setup_browser:
rc = _run_setup_browser(assume_yes=args.assume_yes)
if rc != 0:

View File

@@ -7,7 +7,6 @@ thread-safely onto the loop.
"""
import asyncio
import json
import logging
from collections import deque
from typing import Any, Callable, Deque, Dict
@@ -15,7 +14,7 @@ from typing import Any, Callable, Deque, Dict
import acp
from acp.schema import AgentPlanUpdate, PlanEntry
from .tools import _json_loads_maybe, build_tool_complete, build_tool_start, make_tool_call_id
from .tools import _json_loads_maybe, build_tool_complete, build_tool_start, coerce_tool_args, make_tool_call_id
logger = logging.getLogger(__name__)
@@ -89,14 +88,7 @@ def make_tool_progress_cb(
def _tool_progress(event_type: str, name: str = None, preview: str = None, args: Any = None, **kwargs) -> None:
if event_type != "tool.started":
return
if isinstance(args, str):
try:
args = json.loads(args)
except (json.JSONDecodeError, TypeError):
args = {"raw": args}
if not isinstance(args, dict):
args = {}
args = coerce_tool_args(args)
tc_id = make_tool_call_id()
queue = _upgrade_queue(tool_call_ids, name)
if queue is None:

View File

@@ -6,7 +6,6 @@ import asyncio
from datetime import datetime, timezone
import contextlib
import contextvars
import json
import logging
import os
from collections import Counter, defaultdict, deque
@@ -37,7 +36,7 @@ from acp_adapter.model_catalog import ( # noqa: F401 (ACP_MAX_MODELS_PER_PROVI
from acp_adapter.permissions import make_approval_callback
from acp_adapter.provenance import session_provenance_meta
from acp_adapter.session import SessionManager, SessionState, _expand_acp_enabled_toolsets
from acp_adapter.tools import build_tool_complete, build_tool_start
from acp_adapter.tools import build_tool_complete, build_tool_start, coerce_tool_args
from agent.context_compressor import (COMPRESSED_SUMMARY_METADATA_KEY, ContextCompressor)
from agent.interrupt_compat import request_hard_interrupt
from tools.approval import (reset_hermes_interactive_context, set_hermes_interactive_context)
@@ -116,36 +115,18 @@ def _history_summary_meta(message: dict[str, Any], text: str) -> dict[str, Any]
return None
# role -> (chunk class, session_update tag) for history replay.
_HISTORY_CHUNK_TYPES = {
"user": (UserMessageChunk, "user_message_chunk"),
"assistant": (AgentMessageChunk, "agent_message_chunk"),
}
def _history_message_update(
*, role: str, text: str, field_meta: dict[str, Any] | None = None
) -> UserMessageChunk | AgentMessageChunk | None:
"""ACP history replay update for a user/assistant message."""
spec = _HISTORY_CHUNK_TYPES.get(role)
if spec is None:
return None
cls, session_update = spec
return cls(session_update=session_update, content=TextContentBlock(type="text", text=text), field_meta=field_meta)
def _history_tool_call_name_args(tool_call: dict[str, Any]) -> tuple[str, dict[str, Any]]:
"""Extract function name/arguments from an OpenAI-style tool_call."""
function = tool_call.get("function") if isinstance(tool_call.get("function"), dict) else {}
name = str(function.get("name") or tool_call.get("name") or "unknown_tool")
raw_args = function.get("arguments") or tool_call.get("arguments") or tool_call.get("args") or {}
if isinstance(raw_args, str):
try:
raw_args = json.loads(raw_args)
except Exception:
raw_args = {"raw": raw_args}
if not isinstance(raw_args, dict):
raw_args = {}
return name, raw_args
return name, coerce_tool_args(function.get("arguments") or tool_call.get("arguments") or tool_call.get("args") or {})
def _mcp_server_config(server: McpServerStdio | McpServerHttp | McpServerSse) -> dict:
@@ -183,6 +164,12 @@ def _attach_interrupted_prompt(interrupted_prompt: str, guidance: str) -> str:
return f"{interrupted_prompt}\n\nUser correction/guidance after interrupt: {guidance}"
def _queue_prompt(state: SessionState, text: str) -> int:
with state.runtime_lock:
state.queued_prompts.append(text)
return len(state.queued_prompts)
def _take_interrupted_prompt(state: SessionState) -> tuple[bool, str]:
"""``(idle, interrupted_prompt)``; consumes the cancelled prompt only when the session is idle."""
with state.runtime_lock:
@@ -589,10 +576,11 @@ class HermesACPAgent(acp.Agent):
text = _flatten_history_text(message.get("content"))
if not text:
return True
update = _history_message_update(
role=role, text=text, field_meta=_history_summary_meta(message, text)
)
return update is None or await send(update)
cls, session_update = _HISTORY_CHUNK_TYPES[role]
return await send(cls(
session_update=session_update, content=TextContentBlock(type="text", text=text),
field_meta=_history_summary_meta(message, text),
))
for message in state.history:
role = str(message.get("role") or "")
@@ -702,19 +690,20 @@ class HermesACPAgent(acp.Agent):
async def cancel(self, session_id: str, **kwargs: Any) -> None:
state = self.session_manager.get_session(session_id)
if state and state.cancel_event:
with state.runtime_lock:
if state.is_running and state.current_prompt_text:
state.interrupted_prompt_text = state.current_prompt_text
# Cancel + hard-stop under the lock so no other prompt mistakes this turn for
# redirectable work.
state.cancel_event.set()
try:
if state.agent:
request_hard_interrupt(state.agent)
except Exception:
logger.debug("Failed to interrupt ACP session %s", session_id, exc_info=True)
logger.info("Cancelled session %s", session_id)
if not (state and state.cancel_event):
return
with state.runtime_lock:
if state.is_running and state.current_prompt_text:
state.interrupted_prompt_text = state.current_prompt_text
# Cancel + hard-stop under the lock so no other prompt mistakes this turn for
# redirectable work.
state.cancel_event.set()
try:
if state.agent:
request_hard_interrupt(state.agent)
except Exception:
logger.debug("Failed to interrupt ACP session %s", session_id, exc_info=True)
logger.info("Cancelled session %s", session_id)
async def fork_session(
self, cwd: str, session_id: str, mcp_servers: list | None = None, **kwargs: Any
@@ -817,9 +806,7 @@ class HermesACPAgent(acp.Agent):
if redirected:
return "Redirected the active turn with your correction."
if queued_depth is not None:
return f"Queued for the next turn. ({queued_depth} queued)"
return None
return None if queued_depth is None else f"Queued for the next turn. ({queued_depth} queued)"
def _run_agent_turn(
self, *, state: SessionState, session_id: str, user_text: str, user_content: Any, conn: Any,
@@ -1054,9 +1041,7 @@ class HermesACPAgent(acp.Agent):
total_tokens=result.get("total_tokens", 0), thought_tokens=result.get("reasoning_tokens"),
cached_read_tokens=result.get("cache_read_tokens"),
)
await self._send_usage_update(state)
return PromptResponse(stop_reason="cancelled" if cancelled else "end_turn", usage=usage)
# ---- Slash commands (headless) -------------------------------------------
@@ -1183,12 +1168,11 @@ class HermesACPAgent(acp.Agent):
lines.append(f"Model: {model}")
lines.append(f"Provider: {provider}")
if approx_tokens > 0:
if context_length > 0:
usage_pct = (approx_tokens / context_length) * 100
lines.append(f"Context usage: ~{approx_tokens:,} / {context_length:,} tokens ({usage_pct:.1f}%)")
else:
lines.append(f"Context usage: ~{approx_tokens:,} tokens")
if approx_tokens > 0 and context_length > 0:
usage_pct = (approx_tokens / context_length) * 100
lines.append(f"Context usage: ~{approx_tokens:,} / {context_length:,} tokens ({usage_pct:.1f}%)")
elif approx_tokens > 0:
lines.append(f"Context usage: ~{approx_tokens:,} tokens")
if threshold_tokens > 0:
if approx_tokens > 0:
@@ -1217,18 +1201,15 @@ class HermesACPAgent(acp.Agent):
def _cmd_reset(self, args: str, state: SessionState) -> str:
state.history.clear()
reset_failed = False
try:
reset_session_state = getattr(state.agent, "reset_session_state", None)
if callable(reset_session_state):
reset_session_state()
except Exception:
reset_failed = True
logger.warning("ACP session state reset failed for %s", state.session_id, exc_info=True)
return "Conversation history cleared. Agent session state reset failed; see logs."
finally:
self.session_manager.save_session(state.session_id)
if reset_failed:
return "Conversation history cleared. Agent session state reset failed; see logs."
return "Conversation history cleared."
def _cmd_compress(self, args: str, state: SessionState) -> str:
@@ -1270,11 +1251,6 @@ class HermesACPAgent(acp.Agent):
except Exception as e:
return f"Compression failed: {e}"
def _queue_prompt(self, state: SessionState, text: str) -> int:
with state.runtime_lock:
state.queued_prompts.append(text)
return len(state.queued_prompts)
def _cmd_steer(self, args: str, state: SessionState) -> str:
steer_text = args.strip()
if not steer_text:
@@ -1289,15 +1265,13 @@ class HermesACPAgent(acp.Agent):
logger.warning("ACP steer failed for session %s: %s", state.session_id, exc)
return f"⚠️ Steer failed: {exc}"
depth = self._queue_prompt(state, steer_text)
return f"No active turn — queued for the next turn. ({depth} queued)"
return f"No active turn — queued for the next turn. ({_queue_prompt(state, steer_text)} queued)"
def _cmd_queue(self, args: str, state: SessionState) -> str:
queued_text = args.strip()
if not queued_text:
return "Usage: /queue <prompt>"
depth = self._queue_prompt(state, queued_text)
return f"Queued for the next turn. ({depth} queued)"
return f"Queued for the next turn. ({_queue_prompt(state, queued_text)} queued)"
def _cmd_version(self, args: str, state: SessionState) -> str:
return f"Hermes Agent v{HERMES_VERSION}"

View File

@@ -122,6 +122,17 @@ def _structured(text_fallback: bool = False):
return deco
def coerce_tool_args(raw: Any) -> Args:
"""Tool-call arguments as a dict: JSON strings are decoded (undecodable -> ``{"raw": ...}``),
anything else non-dict becomes ``{}``."""
if isinstance(raw, str):
try:
raw = json.loads(raw)
except Exception:
raw = {"raw": raw}
return raw if isinstance(raw, dict) else {}
def _args_json(arguments: Any) -> str:
try:
return json.dumps(arguments, indent=2, default=str)