refactor(tools): homeassistant/feishu/drive_preview — lark_call helper, merged paging queries, compact registrations
This commit is contained in:
@@ -1,19 +1,13 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Interact with the in-app browser / preview pane in the Hermes desktop GUI.
|
||||
|
||||
``open_preview`` shows a page, ``read_preview`` reads it; this is the third leg —
|
||||
clicking, typing, scrolling, history — so the agent drives the page the user sees.
|
||||
Elements are addressed by legible refs from ``action="elements"`` (``btn-sign-in``,
|
||||
``inp-email``). A ref lasts while the page is open — including across a re-render that
|
||||
rebuilds the element — and only a navigation retires it (the renderer says so rather
|
||||
than acting on whatever now occupies the spot). Because refs hold, the renderer answers
|
||||
with a *delta* (appeared/went/changed/rebound) instead of re-sending the inventory —
|
||||
cheap only because the refs stay legible on their own several turns later.
|
||||
Round-trips through the gateway's blocking-prompt bridge like ``read_preview``:
|
||||
tui_gateway emits ``preview.act.request``, the renderer injects the interaction engine
|
||||
into the pane's webview and answers ``preview.act.respond``. This module is schema + a
|
||||
thin dispatcher over the platform-injected callback. Lives in the ``desktop_ui`` toolset,
|
||||
which the GUI gateway enables only for desktop-sourced sessions.
|
||||
``open_preview`` shows a page, ``read_preview`` reads it; this is the third leg — clicking,
|
||||
typing, scrolling, history. Elements are addressed by legible refs from ``action="elements"``
|
||||
(``btn-sign-in``); a ref survives re-renders and only a navigation retires it, so the renderer
|
||||
answers with a *delta* (appeared/went/changed/rebound) instead of re-sending the inventory.
|
||||
Round-trips through the gateway's blocking-prompt bridge (``preview.act.request`` /
|
||||
``preview.act.respond``); this module is schema + a thin dispatcher over the platform-injected
|
||||
callback. Lives in the ``desktop_ui`` toolset (GUI gateway, desktop-sourced sessions only).
|
||||
"""
|
||||
|
||||
from typing import Callable, Optional
|
||||
@@ -21,18 +15,7 @@ from typing import Callable, Optional
|
||||
from tools.desktop_ui import passthrough_json
|
||||
from tools.registry import registry, tool_error
|
||||
|
||||
ACTIONS = (
|
||||
"elements",
|
||||
"click",
|
||||
"hover",
|
||||
"type",
|
||||
"scroll",
|
||||
"press",
|
||||
"strobe",
|
||||
"back",
|
||||
"forward",
|
||||
"reload",
|
||||
)
|
||||
ACTIONS = ("elements", "click", "hover", "type", "scroll", "press", "strobe", "back", "forward", "reload")
|
||||
SCROLL_TO = ("top", "bottom")
|
||||
|
||||
# Verbs that need something to act on — a ref from the last inventory, or a
|
||||
@@ -56,55 +39,31 @@ def drive_preview_tool(
|
||||
"""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.")
|
||||
|
||||
verb = (action or "").strip().lower()
|
||||
if verb not in ACTIONS:
|
||||
return tool_error(f"action must be one of: {', '.join(ACTIONS)}.")
|
||||
|
||||
if verb in NEEDS_TARGET and not (ref or selector):
|
||||
return tool_error(
|
||||
f"{verb} needs a ref from action='elements' (e.g. 'btn-sign-in') or a CSS selector."
|
||||
)
|
||||
|
||||
return tool_error(f"{verb} needs a ref from action='elements' (e.g. 'btn-sign-in') or a CSS selector.")
|
||||
if verb == "type" and text is None:
|
||||
return tool_error("type needs the text to enter.")
|
||||
|
||||
if verb == "press" and not key:
|
||||
return tool_error("press needs a key, e.g. 'Enter' or 'Escape'.")
|
||||
|
||||
if to is not None and to not in SCROLL_TO:
|
||||
return tool_error(f"to must be one of: {', '.join(SCROLL_TO)}.")
|
||||
|
||||
try:
|
||||
payload = {
|
||||
name: val
|
||||
for name, val in (
|
||||
("action", verb),
|
||||
("ref", ref),
|
||||
("selector", selector),
|
||||
("text", text),
|
||||
("key", key),
|
||||
("submit", submit),
|
||||
("full", full),
|
||||
("to", to),
|
||||
("amount", None if amount is None else int(amount)),
|
||||
("max", None if limit is None else int(limit)),
|
||||
)
|
||||
if val is not None
|
||||
}
|
||||
fields = (
|
||||
("action", verb), ("ref", ref), ("selector", selector), ("text", text), ("key", key),
|
||||
("submit", submit), ("full", full), ("to", to),
|
||||
("amount", None if amount is None else int(amount)), ("max", None if limit is None else int(limit)),
|
||||
)
|
||||
except (TypeError, ValueError):
|
||||
return tool_error("amount and max must be integers.")
|
||||
|
||||
try:
|
||||
raw = callback(payload)
|
||||
raw = callback({name: val for name, val in fields if val is not None})
|
||||
except Exception as exc:
|
||||
return tool_error(f"Failed to act on the in-app browser: {exc}")
|
||||
|
||||
if not raw:
|
||||
return tool_error(
|
||||
"The action timed out, or no GUI window answered. "
|
||||
"Open a page with open_preview first."
|
||||
)
|
||||
return tool_error("The action timed out, or no GUI window answered. Open a page with open_preview first.")
|
||||
return passthrough_json(raw)
|
||||
|
||||
|
||||
@@ -186,17 +145,8 @@ registry.register(
|
||||
toolset="desktop_ui",
|
||||
schema=ACT_PREVIEW_SCHEMA,
|
||||
handler=lambda args, **kw: drive_preview_tool(
|
||||
action=args.get("action", ""),
|
||||
ref=args.get("ref"),
|
||||
selector=args.get("selector"),
|
||||
text=args.get("text"),
|
||||
key=args.get("key"),
|
||||
submit=args.get("submit"),
|
||||
amount=args.get("amount"),
|
||||
to=args.get("to"),
|
||||
limit=args.get("max"),
|
||||
full=args.get("full"),
|
||||
callback=kw.get("callback"),
|
||||
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="🖱️",
|
||||
)
|
||||
|
||||
@@ -5,8 +5,8 @@ Uses the same lazy-import + BaseRequest pattern as feishu_comment.py.
|
||||
"""
|
||||
|
||||
from tools.feishu_lark import ( # noqa: F401 (set_client/get_client are imported by feishu_comment)
|
||||
build_request,
|
||||
_check_feishu,
|
||||
build_request,
|
||||
get_client,
|
||||
raw_body,
|
||||
set_client,
|
||||
@@ -38,24 +38,19 @@ def _handle_feishu_doc_read(args: dict, **kwargs) -> str:
|
||||
doc_token = args.get("doc_token", "").strip()
|
||||
if not doc_token:
|
||||
return tool_error("doc_token is required")
|
||||
|
||||
client = get_client()
|
||||
if client is None:
|
||||
return tool_error("Feishu client not available (not in a Feishu comment context)")
|
||||
|
||||
try:
|
||||
request = build_request("GET", _RAW_CONTENT_URI, paths={"document_id": doc_token})
|
||||
except ImportError:
|
||||
return tool_error("lark_oapi not installed")
|
||||
|
||||
# Tool handlers run synchronously in a worker thread (no running event
|
||||
# loop), so call the blocking lark client directly.
|
||||
# Handlers run synchronously in a worker thread (no event loop): call the blocking client.
|
||||
response = client.request(request)
|
||||
|
||||
code = getattr(response, "code", None)
|
||||
if code != 0:
|
||||
msg = getattr(response, "msg", "unknown error")
|
||||
return tool_error(f"Failed to read document: code={code} msg={msg}")
|
||||
return tool_error(f"Failed to read document: code={code} msg={getattr(response, 'msg', 'unknown error')}")
|
||||
|
||||
body = raw_body(response)
|
||||
if body is not None:
|
||||
@@ -63,24 +58,15 @@ def _handle_feishu_doc_read(args: dict, **kwargs) -> str:
|
||||
return tool_result(success=True, content=body.get("data", {}).get("content", ""))
|
||||
except AttributeError:
|
||||
pass
|
||||
|
||||
# Fallback: try the typed response.data
|
||||
data = getattr(response, "data", None)
|
||||
data = getattr(response, "data", None) # fallback: the typed response.data
|
||||
if data:
|
||||
content = data.get("content", "") if isinstance(data, dict) else getattr(data, "content", str(data))
|
||||
return tool_result(success=True, content=content)
|
||||
|
||||
return tool_error("No content returned from document API")
|
||||
|
||||
|
||||
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",
|
||||
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",
|
||||
)
|
||||
|
||||
@@ -8,10 +8,10 @@ The lark client is injected per-thread by the feishu_comment event handler.
|
||||
import logging
|
||||
|
||||
from tools.feishu_lark import ( # noqa: F401 (set_client/get_client are imported by feishu_comment)
|
||||
build_request,
|
||||
_check_feishu,
|
||||
build_request,
|
||||
get_client,
|
||||
response_data,
|
||||
lark_call,
|
||||
set_client,
|
||||
)
|
||||
from tools.registry import registry, tool_error, tool_result
|
||||
@@ -19,14 +19,6 @@ from tools.registry import registry, tool_error, tool_result
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _do_request(client, method, uri, paths=None, queries=None, body=None):
|
||||
"""Build and execute a BaseRequest, return (code, msg, data_dict)."""
|
||||
# Tool handlers run synchronously in a worker thread (no running event
|
||||
# loop), so call the blocking lark client directly.
|
||||
response = client.request(build_request(method, uri, paths, queries, body))
|
||||
return getattr(response, "code", None), getattr(response, "msg", ""), response_data(response)
|
||||
|
||||
|
||||
def _prepare(args: dict, keys: tuple, missing_msg: str):
|
||||
"""Client check first, then required fields (stripped). Returns (client, values, error|None)."""
|
||||
client = get_client()
|
||||
@@ -42,17 +34,12 @@ def _file_type(args: dict) -> str:
|
||||
return args.get("file_type", "docx") or "docx"
|
||||
|
||||
|
||||
def _paged_queries(args: dict) -> list:
|
||||
"""Query params shared by the comment/reply listing endpoints."""
|
||||
return [
|
||||
("file_type", _file_type(args)),
|
||||
("user_id_type", "open_id"),
|
||||
("page_size", str(args.get("page_size", 100))),
|
||||
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,
|
||||
]
|
||||
|
||||
|
||||
def _with_page_token(queries: list, args: dict) -> list:
|
||||
"""Append page_token last (after any is_whole) so the query order stays as before."""
|
||||
page_token = args.get("page_token", "")
|
||||
if page_token:
|
||||
queries.append(("page_token", page_token))
|
||||
@@ -67,10 +54,6 @@ _REPLIES_URI = "/open-apis/drive/v1/files/:file_token/comments/:comment_id/repli
|
||||
_ADD_COMMENT_URI = "/open-apis/drive/v1/files/:file_token/new_comments"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# feishu_drive_list_comments
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
FEISHU_DRIVE_LIST_COMMENTS_SCHEMA = {
|
||||
"name": "feishu_drive_list_comments",
|
||||
"description": (
|
||||
@@ -103,24 +86,14 @@ def _handle_list_comments(args: dict, **kwargs) -> str:
|
||||
client, (file_token,), err = _prepare(args, ("file_token",), "file_token is required")
|
||||
if err:
|
||||
return err
|
||||
|
||||
queries = _paged_queries(args)
|
||||
if args.get("is_whole", False):
|
||||
queries.append(("is_whole", "true"))
|
||||
_with_page_token(queries, args)
|
||||
|
||||
code, msg, data = _do_request(
|
||||
client, "GET", _COMMENTS_URI, paths={"file_token": file_token}, queries=queries,
|
||||
)
|
||||
extra = (("is_whole", "true"),) if args.get("is_whole", False) else ()
|
||||
code, msg, data = lark_call(
|
||||
client, "GET", _COMMENTS_URI, paths={"file_token": file_token}, queries=_paged_queries(args, *extra))
|
||||
if code != 0:
|
||||
return tool_error(f"List comments failed: code={code} msg={msg}")
|
||||
return tool_result(data)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# feishu_drive_list_comment_replies
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
FEISHU_DRIVE_LIST_REPLIES_SCHEMA = {
|
||||
"name": "feishu_drive_list_comment_replies",
|
||||
"description": "List all replies in a comment thread on a Feishu document.",
|
||||
@@ -147,25 +120,18 @@ FEISHU_DRIVE_LIST_REPLIES_SCHEMA = {
|
||||
|
||||
def _handle_list_replies(args: dict, **kwargs) -> str:
|
||||
client, (file_token, comment_id), err = _prepare(
|
||||
args, ("file_token", "comment_id"), "file_token and comment_id are required"
|
||||
)
|
||||
args, ("file_token", "comment_id"), "file_token and comment_id are required")
|
||||
if err:
|
||||
return err
|
||||
|
||||
code, msg, data = _do_request(
|
||||
client, "GET", _REPLIES_URI,
|
||||
paths={"file_token": file_token, "comment_id": comment_id},
|
||||
queries=_with_page_token(_paged_queries(args), args),
|
||||
code, msg, data = lark_call(
|
||||
client, "GET", _REPLIES_URI, paths={"file_token": file_token, "comment_id": comment_id},
|
||||
queries=_paged_queries(args),
|
||||
)
|
||||
if code != 0:
|
||||
return tool_error(f"List replies failed: code={code} msg={msg}")
|
||||
return tool_result(data)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# feishu_drive_reply_comment
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
FEISHU_DRIVE_REPLY_SCHEMA = {
|
||||
"name": "feishu_drive_reply_comment",
|
||||
"description": (
|
||||
@@ -194,15 +160,12 @@ FEISHU_DRIVE_REPLY_SCHEMA = {
|
||||
|
||||
def _handle_reply_comment(args: dict, **kwargs) -> str:
|
||||
client, (file_token, comment_id, content), err = _prepare(
|
||||
args, ("file_token", "comment_id", "content"), "file_token, comment_id, and content are required"
|
||||
)
|
||||
args, ("file_token", "comment_id", "content"), "file_token, comment_id, and content are required")
|
||||
if err:
|
||||
return err
|
||||
|
||||
# Replies use the rich "content.elements[text_run]" body shape; file_type is a query param.
|
||||
code, msg, data = _do_request(
|
||||
client, "POST", _REPLIES_URI,
|
||||
paths={"file_token": file_token, "comment_id": comment_id},
|
||||
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}}]}},
|
||||
)
|
||||
@@ -211,10 +174,6 @@ def _handle_reply_comment(args: dict, **kwargs) -> str:
|
||||
return tool_result(success=True, data=data)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# feishu_drive_add_comment
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
FEISHU_DRIVE_ADD_COMMENT_SCHEMA = {
|
||||
"name": "feishu_drive_add_comment",
|
||||
"description": (
|
||||
@@ -239,15 +198,12 @@ FEISHU_DRIVE_ADD_COMMENT_SCHEMA = {
|
||||
|
||||
def _handle_add_comment(args: dict, **kwargs) -> str:
|
||||
client, (file_token, content), err = _prepare(
|
||||
args, ("file_token", "content"), "file_token and content are required"
|
||||
)
|
||||
args, ("file_token", "content"), "file_token and content are required")
|
||||
if err:
|
||||
return err
|
||||
|
||||
# new_comments takes the flat "reply_elements[text]" shape with file_type in the body.
|
||||
code, msg, data = _do_request(
|
||||
client, "POST", _ADD_COMMENT_URI,
|
||||
paths={"file_token": file_token},
|
||||
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}]},
|
||||
)
|
||||
if code != 0:
|
||||
@@ -255,28 +211,13 @@ def _handle_add_comment(args: dict, **kwargs) -> str:
|
||||
return tool_result(success=True, data=data)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Registration
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
for _name, _schema, _handler, _desc, _emoji in (
|
||||
("feishu_drive_list_comments", FEISHU_DRIVE_LIST_COMMENTS_SCHEMA, _handle_list_comments,
|
||||
"List document comments", "\U0001f4ac"),
|
||||
("feishu_drive_list_comment_replies", FEISHU_DRIVE_LIST_REPLIES_SCHEMA, _handle_list_replies,
|
||||
"List comment replies", "\U0001f4ac"),
|
||||
("feishu_drive_reply_comment", FEISHU_DRIVE_REPLY_SCHEMA, _handle_reply_comment,
|
||||
"Reply to a document comment", "\u2709\ufe0f"),
|
||||
("feishu_drive_add_comment", FEISHU_DRIVE_ADD_COMMENT_SCHEMA, _handle_add_comment,
|
||||
"Add a whole-document comment", "\u2709\ufe0f"),
|
||||
for _schema, _handler, _desc, _emoji in (
|
||||
(FEISHU_DRIVE_LIST_COMMENTS_SCHEMA, _handle_list_comments, "List document comments", "\U0001f4ac"),
|
||||
(FEISHU_DRIVE_LIST_REPLIES_SCHEMA, _handle_list_replies, "List comment replies", "\U0001f4ac"),
|
||||
(FEISHU_DRIVE_REPLY_SCHEMA, _handle_reply_comment, "Reply to a document comment", "\u2709\ufe0f"),
|
||||
(FEISHU_DRIVE_ADD_COMMENT_SCHEMA, _handle_add_comment, "Add a whole-document comment", "\u2709\ufe0f"),
|
||||
):
|
||||
registry.register(
|
||||
name=_name,
|
||||
toolset="feishu_drive",
|
||||
schema=_schema,
|
||||
handler=_handler,
|
||||
check_fn=_check_feishu,
|
||||
requires_env=[],
|
||||
is_async=False,
|
||||
description=_desc,
|
||||
emoji=_emoji,
|
||||
name=_schema["name"], toolset="feishu_drive", schema=_schema, handler=_handler, check_fn=_check_feishu,
|
||||
requires_env=[], is_async=False, description=_desc, emoji=_emoji,
|
||||
)
|
||||
|
||||
@@ -55,6 +55,16 @@ def build_request(method, uri, paths=None, queries=None, body=None):
|
||||
return builder.build()
|
||||
|
||||
|
||||
def lark_call(client, method, uri, paths=None, queries=None, body=None):
|
||||
"""Build + execute a BaseRequest; returns (code, msg, data_dict).
|
||||
|
||||
Tool handlers run synchronously in a worker thread (no running event loop), so the
|
||||
blocking lark client is called directly.
|
||||
"""
|
||||
response = client.request(build_request(method, uri, paths, queries, body))
|
||||
return getattr(response, "code", None), getattr(response, "msg", ""), response_data(response)
|
||||
|
||||
|
||||
def raw_body(response):
|
||||
"""Parsed JSON object of the raw HTTP body, or None when absent/unparseable/not a dict."""
|
||||
raw = getattr(response, "raw", None)
|
||||
|
||||
@@ -28,6 +28,7 @@ def _get_config():
|
||||
_HASS_TOKEN or get_secret("HASS_TOKEN", "") or "",
|
||||
)
|
||||
|
||||
|
||||
# Valid HA entity_id (e.g. "light.living_room", "sensor.temperature_1").
|
||||
_ENTITY_ID_RE = re.compile(r"^[a-z_][a-z0-9_]*\.[a-z0-9_]+$")
|
||||
|
||||
@@ -52,10 +53,7 @@ def _get_headers(token: str = "") -> Dict[str, str]:
|
||||
"""Return authorization headers for HA REST API."""
|
||||
if not token:
|
||||
_, token = _get_config()
|
||||
return {
|
||||
"Authorization": f"Bearer {token}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
return {"Authorization": f"Bearer {token}", "Content-Type": "application/json"}
|
||||
|
||||
|
||||
async def _api_json(method: str, path: str, timeout: float, payload: Any = None) -> Any:
|
||||
@@ -72,19 +70,12 @@ async def _api_json(method: str, path: str, timeout: float, payload: Any = None)
|
||||
return await resp.json()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Async helpers (called from sync handlers via _run_async)
|
||||
# ---------------------------------------------------------------------------
|
||||
# ── async helpers (called from sync handlers via _run_async) ─────────────────
|
||||
|
||||
def _filter_and_summarize(
|
||||
states: list,
|
||||
domain: Optional[str] = None,
|
||||
area: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Filter raw HA states by domain/area and return a compact summary."""
|
||||
def _filter_and_summarize(states: list, domain: Optional[str] = None, area: Optional[str] = None) -> Dict[str, Any]:
|
||||
"""Filter raw HA states by domain/area (area matches friendly_name or area attr) and compact them."""
|
||||
if domain:
|
||||
states = [s for s in states if s.get("entity_id", "").startswith(f"{domain}.")]
|
||||
|
||||
if area:
|
||||
area_lower = area.lower()
|
||||
states = [
|
||||
@@ -92,11 +83,9 @@ def _filter_and_summarize(
|
||||
if area_lower in (s.get("attributes", {}).get("friendly_name", "") or "").lower()
|
||||
or area_lower in (s.get("attributes", {}).get("area", "") or "").lower()
|
||||
]
|
||||
|
||||
entities = [
|
||||
{
|
||||
"entity_id": s["entity_id"],
|
||||
"state": s["state"],
|
||||
"entity_id": s["entity_id"], "state": s["state"],
|
||||
"friendly_name": s.get("attributes", {}).get("friendly_name", ""),
|
||||
}
|
||||
for s in states
|
||||
@@ -104,17 +93,11 @@ def _filter_and_summarize(
|
||||
return {"count": len(entities), "entities": entities}
|
||||
|
||||
|
||||
async def _async_list_entities(
|
||||
domain: Optional[str] = None,
|
||||
area: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Fetch entity states from HA and optionally filter by domain/area."""
|
||||
states = await _api_json("GET", "/api/states", 15)
|
||||
return _filter_and_summarize(states, domain, area)
|
||||
async def _async_list_entities(domain: Optional[str] = None, area: Optional[str] = None) -> Dict[str, Any]:
|
||||
return _filter_and_summarize(await _api_json("GET", "/api/states", 15), domain, area)
|
||||
|
||||
|
||||
async def _async_get_state(entity_id: str) -> Dict[str, Any]:
|
||||
"""Fetch detailed state of a single entity."""
|
||||
data = await _api_json("GET", f"/api/states/{entity_id}", 10)
|
||||
return {
|
||||
"entity_id": data["entity_id"],
|
||||
@@ -125,52 +108,34 @@ async def _async_get_state(entity_id: str) -> Dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _build_service_payload(
|
||||
entity_id: Optional[str] = None,
|
||||
data: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Build the JSON payload for a HA service call; ``entity_id`` overrides data["entity_id"]."""
|
||||
def _build_service_payload(entity_id: Optional[str] = None, data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||||
"""JSON payload for a HA service call; ``entity_id`` overrides data["entity_id"]."""
|
||||
payload: Dict[str, Any] = dict(data or {})
|
||||
if entity_id:
|
||||
payload["entity_id"] = entity_id
|
||||
return payload
|
||||
|
||||
|
||||
def _parse_service_response(
|
||||
domain: str,
|
||||
service: str,
|
||||
result: Any,
|
||||
) -> Dict[str, Any]:
|
||||
"""Parse HA service call response into a structured result."""
|
||||
def _parse_service_response(domain: str, service: str, result: Any) -> Dict[str, Any]:
|
||||
affected = []
|
||||
if isinstance(result, list):
|
||||
affected = [{"entity_id": s.get("entity_id", ""), "state": s.get("state", "")} for s in result]
|
||||
return {
|
||||
"success": True,
|
||||
"service": f"{domain}.{service}",
|
||||
"affected_entities": affected,
|
||||
}
|
||||
return {"success": True, "service": f"{domain}.{service}", "affected_entities": affected}
|
||||
|
||||
|
||||
async def _async_call_service(
|
||||
domain: str,
|
||||
service: str,
|
||||
entity_id: Optional[str] = None,
|
||||
data: Optional[Dict[str, Any]] = None,
|
||||
domain: str, service: str, entity_id: Optional[str] = None, data: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Call a Home Assistant service."""
|
||||
result = await _api_json(
|
||||
"POST", f"/api/services/{domain}/{service}", 15, _build_service_payload(entity_id, data)
|
||||
)
|
||||
"POST", f"/api/services/{domain}/{service}", 15, _build_service_payload(entity_id, data))
|
||||
return _parse_service_response(domain, service, result)
|
||||
|
||||
|
||||
async def _async_list_services(domain: Optional[str] = None) -> Dict[str, Any]:
|
||||
"""Fetch available services from HA, optionally filtered by domain, compacted for context."""
|
||||
"""Available services, optionally filtered by domain, compacted for context."""
|
||||
services = await _api_json("GET", "/api/services", 15)
|
||||
if domain:
|
||||
services = [s for s in services if s.get("domain") == domain]
|
||||
|
||||
result = []
|
||||
for svc_domain in services:
|
||||
domain_services = {}
|
||||
@@ -178,19 +143,13 @@ async def _async_list_services(domain: Optional[str] = None) -> Dict[str, Any]:
|
||||
svc_entry: Dict[str, Any] = {"description": svc_info.get("description", "")}
|
||||
fields = svc_info.get("fields", {})
|
||||
if fields:
|
||||
svc_entry["fields"] = {
|
||||
k: v.get("description", "") for k, v in fields.items()
|
||||
if isinstance(v, dict)
|
||||
}
|
||||
svc_entry["fields"] = {k: v.get("description", "") for k, v in fields.items() if isinstance(v, dict)}
|
||||
domain_services[svc_name] = svc_entry
|
||||
result.append({"domain": svc_domain.get("domain", ""), "services": domain_services})
|
||||
|
||||
return {"count": len(result), "domains": result}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Sync wrappers (handler signature: (args, **kw) -> str)
|
||||
# ---------------------------------------------------------------------------
|
||||
# ── sync wrappers (handler signature: (args, **kw) -> str) ───────────────────
|
||||
|
||||
def _run_async(coro):
|
||||
"""Run a coroutine from a sync handler; hops to a thread if a loop is already running."""
|
||||
@@ -198,7 +157,6 @@ def _run_async(coro):
|
||||
loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
loop = None
|
||||
|
||||
if loop and loop.is_running():
|
||||
import concurrent.futures
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
|
||||
@@ -236,30 +194,25 @@ def _handle_call_service(args: dict, **kw) -> str:
|
||||
service = args.get("service", "")
|
||||
if not domain or not service:
|
||||
return tool_error("Missing required parameters: domain and service")
|
||||
|
||||
# Format check BEFORE the blocklist: rejects "shell_command/../light" style bypasses.
|
||||
if not _SERVICE_NAME_RE.match(domain):
|
||||
return tool_error(f"Invalid domain format: {domain!r}")
|
||||
if not _SERVICE_NAME_RE.match(service):
|
||||
return tool_error(f"Invalid service format: {service!r}")
|
||||
|
||||
if domain in _BLOCKED_DOMAINS:
|
||||
return tool_error(
|
||||
f"Service domain '{domain}' is blocked for security. "
|
||||
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}")
|
||||
|
||||
data = args.get("data")
|
||||
if isinstance(data, str): # XML tool-calling mode delivers data as a JSON string
|
||||
try:
|
||||
data = json.loads(data) if data.strip() else None
|
||||
except json.JSONDecodeError as e:
|
||||
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}",
|
||||
@@ -268,8 +221,7 @@ def _handle_call_service(args: dict, **kw) -> str:
|
||||
|
||||
def _handle_list_services(args: dict, **kw) -> str:
|
||||
return _dispatch(
|
||||
_async_list_services(domain=args.get("domain")), "ha_list_services", "Failed to list services"
|
||||
)
|
||||
_async_list_services(domain=args.get("domain")), "ha_list_services", "Failed to list services")
|
||||
|
||||
|
||||
def _check_ha_available() -> bool:
|
||||
@@ -277,9 +229,7 @@ def _check_ha_available() -> bool:
|
||||
return bool(get_secret("HASS_TOKEN"))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tool schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
# ── tool schemas ─────────────────────────────────────────────────────────────
|
||||
|
||||
HA_LIST_ENTITIES_SCHEMA = {
|
||||
"name": "ha_list_entities",
|
||||
@@ -401,10 +351,6 @@ HA_CALL_SERVICE_SCHEMA = {
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Registration
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
for _schema, _handler in (
|
||||
(HA_LIST_ENTITIES_SCHEMA, _handle_list_entities),
|
||||
(HA_GET_STATE_SCHEMA, _handle_get_state),
|
||||
@@ -412,10 +358,6 @@ for _schema, _handler in (
|
||||
(HA_CALL_SERVICE_SCHEMA, _handle_call_service),
|
||||
):
|
||||
registry.register(
|
||||
name=_schema["name"],
|
||||
toolset="homeassistant",
|
||||
schema=_schema,
|
||||
handler=_handler,
|
||||
check_fn=_check_ha_available,
|
||||
emoji="🏠",
|
||||
name=_schema["name"], toolset="homeassistant", schema=_schema, handler=_handler,
|
||||
check_fn=_check_ha_available, emoji="🏠",
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user