refactor(tools): group I — AST-neutral layout compaction (closers/hug), schemas untouched

This commit is contained in:
Teknium
2026-09-02 22:01:02 -07:00
parent 7ed389e178
commit ca2867073b
6 changed files with 45 additions and 92 deletions

View File

@@ -56,8 +56,7 @@ def _get_bot_token() -> Optional[str]:
def _discord_request(
method: str, path: str, token: str, params: Optional[Dict[str, str]] = None,
body: Optional[Dict[str, Any]] = None, timeout: int = 15,
) -> Any:
body: Optional[Dict[str, Any]] = None, timeout: int = 15) -> Any:
"""Make a request to the Discord REST API."""
url = f"{DISCORD_API_BASE}{path}"
if params:
@@ -67,9 +66,7 @@ def _discord_request(
headers={
"Authorization": f"Bot {token}",
"Content-Type": "application/json",
"User-Agent": "Hermes-Agent (https://github.com/NousResearch/hermes-agent)",
},
)
"User-Agent": "Hermes-Agent (https://github.com/NousResearch/hermes-agent)"})
try:
with urllib.request.urlopen(req, timeout=timeout) as resp:
if resp.status == 204:
@@ -91,8 +88,7 @@ def _discord_request(
_CHANNEL_TYPE_NAMES = {
0: "text", 2: "voice", 4: "category", 5: "announcement", 10: "announcement_thread",
11: "public_thread", 12: "private_thread", 13: "stage", 15: "forum", 16: "media",
}
11: "public_thread", 12: "private_thread", 13: "stage", 15: "forum", 16: "media"}
def _channel_type_name(type_id: int) -> str:
@@ -249,15 +245,13 @@ def _member_summary(m: Dict[str, Any], *, full: bool) -> Dict[str, Any]:
user = m.get("user", {})
base = {
"user_id": user.get("id"), "username": user.get("username"),
"display_name": user.get("global_name"), "nickname": m.get("nick"),
}
"display_name": user.get("global_name"), "nickname": m.get("nick")}
tail = {"bot": user.get("bot", False), "roles": m.get("roles", [])}
if not full:
return {**base, **tail}
return {
**base, "avatar": user.get("avatar"), **tail,
"joined_at": m.get("joined_at"), "premium_since": m.get("premium_since"),
}
"joined_at": m.get("joined_at"), "premium_since": m.get("premium_since")}
def _message_summary(msg: Dict[str, Any]) -> Dict[str, Any]:
@@ -267,20 +261,17 @@ def _message_summary(msg: Dict[str, Any]) -> Dict[str, Any]:
"content": msg.get("content", ""),
"author": {
"id": author.get("id"), "username": author.get("username"),
"display_name": author.get("global_name"), "bot": author.get("bot", False),
},
"display_name": author.get("global_name"), "bot": author.get("bot", False)},
"timestamp": msg.get("timestamp"),
"edited_timestamp": msg.get("edited_timestamp"),
"attachments": [
{"filename": a.get("filename"), "url": a.get("url"), "size": a.get("size")}
for a in msg.get("attachments", [])
],
for a in msg.get("attachments", [])],
"reactions": [
{"emoji": r.get("emoji", {}).get("name"), "count": r.get("count", 0)}
for r in msg.get("reactions", [])
] if msg.get("reactions") else [],
"pinned": msg.get("pinned", False),
}
"pinned": msg.get("pinned", False)}
def _int_or(value: Any, default: int) -> int:
@@ -300,10 +291,8 @@ def _list_guilds(token: str, **_kwargs: Any) -> str:
return _listing("guilds", [
{
"id": g["id"], "name": g["name"], "icon": g.get("icon"),
"owner": g.get("owner", False), "permissions": g.get("permissions"),
}
for g in guilds
])
"owner": g.get("owner", False), "permissions": g.get("permissions")}
for g in guilds])
def _server_info(token: str, guild_id: str, **_kwargs: Any) -> str:
@@ -319,8 +308,7 @@ def _server_info(token: str, guild_id: str, **_kwargs: Any) -> str:
"features": g.get("features", []),
"premium_tier": g.get("premium_tier"),
"premium_subscription_count": g.get("premium_subscription_count"),
"verification_level": g.get("verification_level"),
})
"verification_level": g.get("verification_level")})
def _list_channels(token: str, guild_id: str, **_kwargs: Any) -> str:
@@ -366,8 +354,7 @@ def _channel_info(token: str, channel_id: str, **_kwargs: Any) -> str:
"position": ch.get("position"),
"parent_id": ch.get("parent_id"),
"rate_limit_per_user": ch.get("rate_limit_per_user", 0),
"last_message_id": ch.get("last_message_id"),
})
"last_message_id": ch.get("last_message_id")})
def _list_roles(token: str, guild_id: str, **_kwargs: Any) -> str:
@@ -381,10 +368,8 @@ def _list_roles(token: str, guild_id: str, **_kwargs: Any) -> str:
"mentionable": r.get("mentionable", False),
"managed": r.get("managed", False),
"member_count": r.get("member_count"),
"hoist": r.get("hoist", False),
}
for r in sorted(roles, key=lambda r: r.get("position", 0), reverse=True)
])
"hoist": r.get("hoist", False)}
for r in sorted(roles, key=lambda r: r.get("position", 0), reverse=True)])
def _member_info(token: str, guild_id: str, user_id: str, **_kwargs: Any) -> str:
@@ -401,8 +386,7 @@ def _search_members(token: str, guild_id: str, query: str, limit: int = 20, **_k
def _fetch_messages(
token: str, channel_id: str, limit: int = 50,
before: Optional[str] = None, after: Optional[str] = None, **_kwargs: Any,
) -> str:
before: Optional[str] = None, after: Optional[str] = None, **_kwargs: Any) -> str:
"""``before``/``after`` are message snowflakes for reverse/forward pagination."""
params: Dict[str, str] = {"limit": _limit_param(limit, 50)}
if before:
@@ -419,16 +403,13 @@ def _list_pins(token: str, channel_id: str, **_kwargs: Any) -> str:
return _listing("pinned_messages", [
{
"id": msg["id"], "content": msg.get("content", "")[:200],
"author": msg.get("author", {}).get("username"), "timestamp": msg.get("timestamp"),
}
for msg in messages
])
"author": msg.get("author", {}).get("username"), "timestamp": msg.get("timestamp")}
for msg in messages])
def _create_thread(
token: str, channel_id: str, name: str, message_id: Optional[str] = None,
auto_archive_duration: int = 1440, **_kwargs: Any,
) -> str:
auto_archive_duration: int = 1440, **_kwargs: Any) -> str:
"""Create a thread — anchored to ``message_id`` when given, else standalone public."""
body: Dict[str, Any] = {"name": name, "auto_archive_duration": auto_archive_duration}
if message_id:
@@ -517,8 +498,7 @@ def _load_allowed_actions_config() -> Optional[List[str]]:
if invalid:
logger.warning(
"discord.server_actions: unknown action(s) ignored: %s. Known: %s",
", ".join(invalid), ", ".join(_ACTIONS.keys()),
)
", ".join(invalid), ", ".join(_ACTIONS.keys()))
return [n for n in names if n in _ACTIONS]
@@ -587,8 +567,7 @@ def _build_schema(
if not actions:
return None
manifest_block = "\n".join(
f" {name}{sig} — {desc}" for name, _fn, sig, desc in _ACTION_MANIFEST if name in actions
)
f" {name}{sig} — {desc}" for name, _fn, sig, desc in _ACTION_MANIFEST if name in actions)
content_note = ""
affected_actions = {"fetch_messages", "list_pins"} & set(actions)
if affected_actions and caps.get("detected") and caps.get("has_message_content") is False:
@@ -600,9 +579,7 @@ def _build_schema(
"parameters": {
"type": "object",
"properties": {"action": {"type": "string", "enum": actions}, **_SCHEMA_PROPERTIES},
"required": ["action"],
},
}
"required": ["action"]}}
def _get_dynamic_schema(action_subset: Dict[str, Any], tool_name: str) -> Optional[Dict[str, Any]]:
@@ -677,8 +654,7 @@ def check_discord_tool_requirements() -> bool:
_HANDLER_DEFAULTS = {
"guild_id": "", "channel_id": "", "user_id": "", "role_id": "", "message_id": "", "query": "",
"name": "", "limit": 50, "before": "", "after": "", "auto_archive_duration": 1440,
}
"name": "", "limit": 50, "before": "", "after": "", "auto_archive_duration": 1440}
def _run_discord_action(action: str, valid_actions: Dict[str, Any], tool_label: str, **params: Any) -> str:
@@ -696,8 +672,7 @@ def _run_discord_action(action: str, valid_actions: Dict[str, Any], tool_label:
if allowlist is not None and action not in allowlist:
return tool_error(
f"Action '{action}' is disabled by config (discord.server_actions). "
f"Allowed: {', '.join(allowlist) if allowlist else '<none>'}"
)
f"Allowed: {', '.join(allowlist) if allowlist else '<none>'}")
kwargs = {k: params.get(k, v) for k, v in _HANDLER_DEFAULTS.items()}
missing = [p for p in _REQUIRED_PARAMS.get(action, []) if not kwargs.get(p)]
@@ -736,5 +711,4 @@ for _name, _actions, _handler in (
handler=lambda args, _h=_handler, **kw: _h(
action=args.get("action", ""), **{k: args.get(k, v) for k, v in _HANDLER_DEFAULTS.items()}),
check_fn=check_discord_tool_requirements,
requires_env=["DISCORD_BOT_TOKEN"],
)
requires_env=["DISCORD_BOT_TOKEN"])

View File

@@ -34,8 +34,7 @@ def drive_preview_tool(
to: Optional[str] = None,
limit: Optional[int] = None,
full: Optional[bool] = None,
callback: Optional[Callable] = None,
) -> str:
callback: Optional[Callable] = None) -> str:
"""Dispatch one interaction to the desktop renderer and return its outcome."""
if callback is None:
return tool_error("drive_preview is only available in the Hermes desktop app.")
@@ -148,5 +147,4 @@ registry.register(
action=args.get("action", ""), limit=args.get("max"), callback=kw.get("callback"),
**{k: args.get(k) for k in ("ref", "selector", "text", "key", "submit", "amount", "to", "full")},
),
emoji="🖱️",
)
emoji="🖱️")

View File

@@ -9,8 +9,7 @@ from tools.feishu_lark import ( # noqa: F401 (set_client/get_client are import
build_request,
get_client,
raw_body,
set_client,
)
set_client)
from tools.registry import registry, tool_error, tool_result
_RAW_CONTENT_URI = "/open-apis/docx/v1/documents/:document_id/raw_content"
@@ -68,5 +67,4 @@ def _handle_feishu_doc_read(args: dict, **kwargs) -> str:
registry.register(
name="feishu_doc_read", toolset="feishu_doc", schema=FEISHU_DOC_READ_SCHEMA, handler=_handle_feishu_doc_read,
check_fn=_check_feishu, requires_env=[], is_async=False, description="Read Feishu document content",
emoji="\U0001f4c4",
)
emoji="\U0001f4c4")

View File

@@ -12,8 +12,7 @@ from tools.feishu_lark import ( # noqa: F401 (set_client/get_client are import
build_request,
get_client,
lark_call,
set_client,
)
set_client)
from tools.registry import registry, tool_error, tool_result
logger = logging.getLogger(__name__)
@@ -38,8 +37,7 @@ def _paged_queries(args: dict, *extra) -> list:
"""Query params shared by the listing endpoints; page_token goes last (after any extra)."""
queries = [
("file_type", _file_type(args)), ("user_id_type", "open_id"),
("page_size", str(args.get("page_size", 100))), *extra,
]
("page_size", str(args.get("page_size", 100))), *extra]
page_token = args.get("page_token", "")
if page_token:
queries.append(("page_token", page_token))
@@ -125,8 +123,7 @@ def _handle_list_replies(args: dict, **kwargs) -> str:
return err
code, msg, data = lark_call(
client, "GET", _REPLIES_URI, paths={"file_token": file_token, "comment_id": comment_id},
queries=_paged_queries(args),
)
queries=_paged_queries(args))
if code != 0:
return tool_error(f"List replies failed: code={code} msg={msg}")
return tool_result(data)
@@ -167,8 +164,7 @@ def _handle_reply_comment(args: dict, **kwargs) -> str:
code, msg, data = lark_call(
client, "POST", _REPLIES_URI, paths={"file_token": file_token, "comment_id": comment_id},
queries=[("file_type", _file_type(args))],
body={"content": {"elements": [{"type": "text_run", "text_run": {"text": content}}]}},
)
body={"content": {"elements": [{"type": "text_run", "text_run": {"text": content}}]}})
if code != 0:
return tool_error(f"Reply comment failed: code={code} msg={msg}")
return tool_result(success=True, data=data)
@@ -204,8 +200,7 @@ def _handle_add_comment(args: dict, **kwargs) -> str:
# new_comments takes the flat "reply_elements[text]" shape with file_type in the body.
code, msg, data = lark_call(
client, "POST", _ADD_COMMENT_URI, paths={"file_token": file_token},
body={"file_type": _file_type(args), "reply_elements": [{"type": "text", "text": content}]},
)
body={"file_type": _file_type(args), "reply_elements": [{"type": "text", "text": content}]})
if code != 0:
return tool_error(f"Add comment failed: code={code} msg={msg}")
return tool_result(success=True, data=data)
@@ -219,5 +214,4 @@ for _schema, _handler, _desc, _emoji in (
):
registry.register(
name=_schema["name"], toolset="feishu_drive", schema=_schema, handler=_handler, check_fn=_check_feishu,
requires_env=[], is_async=False, description=_desc, emoji=_emoji,
)
requires_env=[], is_async=False, description=_desc, emoji=_emoji)

View File

@@ -44,8 +44,7 @@ def build_request(method, uri, paths=None, queries=None, body=None):
BaseRequest.builder()
.http_method(HttpMethod.GET if method == "GET" else HttpMethod.POST)
.uri(uri)
.token_types({AccessTokenType.TENANT})
)
.token_types({AccessTokenType.TENANT}))
if paths:
builder = builder.paths(paths)
if queries:

View File

@@ -25,8 +25,7 @@ def _get_config():
"""Return the active profile's Home Assistant URL and token."""
return (
(_HASS_URL or get_secret("HASS_URL", "http://homeassistant.local:8123") or "").rstrip("/"),
_HASS_TOKEN or get_secret("HASS_TOKEN", "") or "",
)
_HASS_TOKEN or get_secret("HASS_TOKEN", "") or "")
# Valid HA entity_id (e.g. "light.living_room", "sensor.temperature_1").
@@ -81,15 +80,12 @@ def _filter_and_summarize(states: list, domain: Optional[str] = None, area: Opti
states = [
s for s in states
if area_lower in (s.get("attributes", {}).get("friendly_name", "") or "").lower()
or area_lower in (s.get("attributes", {}).get("area", "") or "").lower()
]
or area_lower in (s.get("attributes", {}).get("area", "") or "").lower()]
entities = [
{
"entity_id": s["entity_id"], "state": s["state"],
"friendly_name": s.get("attributes", {}).get("friendly_name", ""),
}
for s in states
]
"friendly_name": s.get("attributes", {}).get("friendly_name", "")}
for s in states]
return {"count": len(entities), "entities": entities}
@@ -104,8 +100,7 @@ async def _async_get_state(entity_id: str) -> Dict[str, Any]:
"state": data["state"],
"attributes": data.get("attributes", {}),
"last_changed": data.get("last_changed"),
"last_updated": data.get("last_updated"),
}
"last_updated": data.get("last_updated")}
def _build_service_payload(entity_id: Optional[str] = None, data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
@@ -176,8 +171,7 @@ def _dispatch(coro, log_name: str, fail_msg: str) -> str:
def _handle_list_entities(args: dict, **kw) -> str:
return _dispatch(
_async_list_entities(domain=args.get("domain"), area=args.get("area")),
"ha_list_entities", "Failed to list entities",
)
"ha_list_entities", "Failed to list entities")
def _handle_get_state(args: dict, **kw) -> str:
@@ -202,8 +196,7 @@ def _handle_call_service(args: dict, **kw) -> str:
if domain in _BLOCKED_DOMAINS:
return tool_error(
f"Service domain '{domain}' is blocked for security. "
f"Blocked domains: {', '.join(sorted(_BLOCKED_DOMAINS))}"
)
f"Blocked domains: {', '.join(sorted(_BLOCKED_DOMAINS))}")
entity_id = args.get("entity_id")
if entity_id and not _ENTITY_ID_RE.match(entity_id):
return tool_error(f"Invalid entity_id format: {entity_id}")
@@ -215,8 +208,7 @@ def _handle_call_service(args: dict, **kw) -> str:
return tool_error(f"Invalid JSON string in 'data' parameter: {e}")
return _dispatch(
_async_call_service(domain, service, entity_id, data),
"ha_call_service", f"Failed to call {domain}.{service}",
)
"ha_call_service", f"Failed to call {domain}.{service}")
def _handle_list_services(args: dict, **kw) -> str:
@@ -355,9 +347,7 @@ for _schema, _handler in (
(HA_LIST_ENTITIES_SCHEMA, _handle_list_entities),
(HA_GET_STATE_SCHEMA, _handle_get_state),
(HA_LIST_SERVICES_SCHEMA, _handle_list_services),
(HA_CALL_SERVICE_SCHEMA, _handle_call_service),
):
(HA_CALL_SERVICE_SCHEMA, _handle_call_service)):
registry.register(
name=_schema["name"], toolset="homeassistant", schema=_schema, handler=_handler,
check_fn=_check_ha_available, emoji="🏠",
)
check_fn=_check_ha_available, emoji="🏠")