diff --git a/agent/middleware/__init__.py b/agent/middleware/__init__.py index da26db6b..a1678d97 100644 --- a/agent/middleware/__init__.py +++ b/agent/middleware/__init__.py @@ -2,6 +2,7 @@ from .check_message_queue import check_message_queue_before_model from .ensure_no_empty_msg import ensure_no_empty_msg from .exclude_tools import ExcludeToolsMiddleware from .notify_step_limit import notify_step_limit_reached +from .refresh_slack_status import SlackAssistantStatusMiddleware from .sanitize_tool_inputs import SanitizeToolInputsMiddleware from .tool_error_handler import ToolErrorMiddleware @@ -9,6 +10,7 @@ __all__ = [ "ExcludeToolsMiddleware", "SanitizeToolInputsMiddleware", "ToolErrorMiddleware", + "SlackAssistantStatusMiddleware", "check_message_queue_before_model", "ensure_no_empty_msg", "notify_step_limit_reached", diff --git a/agent/middleware/refresh_slack_status.py b/agent/middleware/refresh_slack_status.py new file mode 100644 index 00000000..f1d8b69e --- /dev/null +++ b/agent/middleware/refresh_slack_status.py @@ -0,0 +1,229 @@ +"""Middleware that keeps Slack's assistant status current during agent work. + +Slack's ``assistants.threads.setStatus`` indicator expires after two minutes +if no message is sent. This middleware refreshes the indicator while model and +tool calls are actively running, then clears it when the run exits without +relying on the model to post a final Slack reply. +""" + +from __future__ import annotations + +import asyncio +import contextlib +import logging +from collections.abc import Awaitable, Callable +from typing import Any, TypeVar + +from langchain.agents.middleware.types import ( + AgentMiddleware, + AgentState, + ModelRequest, + ModelResponse, +) +from langchain_core.messages import ToolMessage +from langgraph.config import get_config +from langgraph.prebuilt.tool_node import ToolCallRequest +from langgraph.runtime import Runtime +from langgraph.types import Command + +from ..utils.slack import ( + DEFAULT_ASSISTANT_STATUS, + DEFAULT_LOADING_MESSAGES, + set_slack_assistant_status, +) + +logger = logging.getLogger(__name__) + +_HEARTBEAT_INTERVAL_SECONDS = 60.0 +_MAX_HEARTBEAT_SECONDS = 60 * 60 + +_T = TypeVar("_T") + + +# Tool-name -> human-readable status. Keep in sync with the tool list in +# agent/server.py, agent/reviewer.py, and the deepagents built-ins (read_file, +# write_file, edit_file, execute, glob, grep, task). +_TOOL_STATUS: dict[str, str] = { + "read_file": "reading files...", + "write_file": "editing files...", + "edit_file": "editing files...", + "execute": "running commands...", + "glob": "scanning the repo...", + "grep": "searching the codebase...", + "task": "delegating to a subagent...", + "web_search": "searching the web...", + "fetch_url": "fetching a URL...", + "http_request": "making an HTTP request...", + "request_pr_review": "requesting a PR review...", + "slack_read_thread_messages": "reading Slack history...", + "slack_thread_reply": "drafting a Slack reply...", + "linear_comment": "commenting on Linear...", + "linear_create_issue": "creating a Linear issue...", + "linear_get_issue": "checking Linear...", + "linear_get_issue_comments": "checking Linear...", + "linear_list_teams": "checking Linear...", + "linear_update_issue": "updating Linear...", + "linear_delete_issue": "updating Linear...", + "add_finding": "recording review findings...", + "update_finding": "updating review findings...", + "list_findings": "checking review findings...", + "publish_review": "publishing the review...", +} + + +def _slack_thread_from_config() -> tuple[str, str] | None: + config = get_config() + configurable = config.get("configurable", {}) if isinstance(config, dict) else {} + slack_thread = configurable.get("slack_thread") if isinstance(configurable, dict) else None + if not isinstance(slack_thread, dict): + return None + + channel_id = slack_thread.get("channel_id") + thread_ts = slack_thread.get("thread_ts") + if not isinstance(channel_id, str) or not isinstance(thread_ts, str): + return None + if not channel_id or not thread_ts: + return None + return channel_id, thread_ts + + +def _tool_call_name(tool_call: object) -> str | None: + if isinstance(tool_call, dict): + name = tool_call.get("name") + else: + name = getattr(tool_call, "name", None) + return name if isinstance(name, str) and name else None + + +def _status_from_recent_tool_calls(messages: list[Any]) -> str: + """Pick a status string based on the last assistant message's tool calls.""" + for msg in reversed(messages): + tool_calls = getattr(msg, "tool_calls", None) + if not tool_calls: + continue + # Use the first tool call's name; if the agent fans out, this is fine + # as a single-line indicator. + name = _tool_call_name(tool_calls[0]) + if isinstance(name, str) and name in _TOOL_STATUS: + return _TOOL_STATUS[name] + return DEFAULT_ASSISTANT_STATUS + return DEFAULT_ASSISTANT_STATUS + + +async def _set_status(channel_id: str, thread_ts: str, status: str) -> None: + try: + await set_slack_assistant_status( + channel_id, + thread_ts, + status=status, + loading_messages=list(DEFAULT_LOADING_MESSAGES) if status else None, + ) + except Exception: + logger.exception("Failed to update Slack assistant status") + + +class SlackAssistantStatusMiddleware(AgentMiddleware): + """Maintain Slack's assistant status for Slack-triggered agent runs.""" + + state_schema = AgentState + + def __init__( + self, + *, + heartbeat_interval_seconds: float = _HEARTBEAT_INTERVAL_SECONDS, + max_heartbeat_seconds: float = _MAX_HEARTBEAT_SECONDS, + ) -> None: + self._heartbeat_interval_seconds = heartbeat_interval_seconds + self._max_heartbeat_seconds = max_heartbeat_seconds + + async def abefore_agent( + self, + state: AgentState, # noqa: ARG002 + runtime: Runtime, # noqa: ARG002 + ) -> dict[str, Any] | None: + await self._try_set(DEFAULT_ASSISTANT_STATUS) + return None + + async def aafter_agent( + self, + state: AgentState, # noqa: ARG002 + runtime: Runtime, # noqa: ARG002 + ) -> dict[str, Any] | None: + # Slack auto-clears on bot replies. This explicit clear covers soft + # exits where the agent stops without posting a message. + await self._try_set("") + return None + + async def awrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], Awaitable[ModelResponse]], + ) -> ModelResponse: + messages = request.state.get("messages", []) if isinstance(request.state, dict) else [] + status = _status_from_recent_tool_calls(messages) + return await self._run_with_heartbeat(status, handler(request)) + + async def awrap_tool_call( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command]], + ) -> ToolMessage | Command: + name = _tool_call_name(request.tool_call) + status = _TOOL_STATUS.get(name or "", DEFAULT_ASSISTANT_STATUS) + return await self._run_with_heartbeat(status, handler(request)) + + def wrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], ModelResponse], + ) -> ModelResponse: + return handler(request) + + def wrap_tool_call( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], ToolMessage | Command], + ) -> ToolMessage | Command: + return handler(request) + + async def _try_set(self, status: str) -> None: + try: + slack_thread = _slack_thread_from_config() + if slack_thread is None: + return + channel_id, thread_ts = slack_thread + await _set_status(channel_id, thread_ts, status) + except Exception: + logger.exception("Failed to read Slack thread config") + + async def _run_with_heartbeat(self, status: str, awaitable: Awaitable[_T]) -> _T: + try: + slack_thread = _slack_thread_from_config() + except Exception: + logger.exception("Failed to read Slack thread config") + return await awaitable + if slack_thread is None: + return await awaitable + + channel_id, thread_ts = slack_thread + await _set_status(channel_id, thread_ts, status) + heartbeat = asyncio.create_task(self._heartbeat(channel_id, thread_ts, status)) + try: + return await awaitable + finally: + heartbeat.cancel() + with contextlib.suppress(asyncio.CancelledError): + await heartbeat + + async def _heartbeat(self, channel_id: str, thread_ts: str, status: str) -> None: + loop = asyncio.get_running_loop() + started_at = loop.time() + while True: + await asyncio.sleep(self._heartbeat_interval_seconds) + if loop.time() - started_at >= self._max_heartbeat_seconds: + logger.info( + "Stopping Slack assistant status heartbeat after %.0f seconds", + self._max_heartbeat_seconds, + ) + return + await _set_status(channel_id, thread_ts, status) diff --git a/agent/reviewer.py b/agent/reviewer.py index cfdcb018..eaf77c2e 100644 --- a/agent/reviewer.py +++ b/agent/reviewer.py @@ -33,6 +33,7 @@ from langchain.agents.middleware import ModelCallLimitMiddleware from .middleware import ( ExcludeToolsMiddleware, SanitizeToolInputsMiddleware, + SlackAssistantStatusMiddleware, ToolErrorMiddleware, ) from .reviewer_findings import ( @@ -347,6 +348,7 @@ async def get_reviewer_agent(config: RunnableConfig) -> Pregel: SanitizeToolInputsMiddleware(), ModelCallLimitMiddleware(run_limit=MODEL_CALL_RECURSION_LIMIT, exit_behavior="end"), ToolErrorMiddleware(), + SlackAssistantStatusMiddleware(), ExcludeToolsMiddleware(excluded=frozenset({"task"})), ], ).with_config(config) diff --git a/agent/server.py b/agent/server.py index 12e6ad65..2a77f518 100644 --- a/agent/server.py +++ b/agent/server.py @@ -30,6 +30,7 @@ from langsmith.sandbox import SandboxClientError from .integrations.langsmith import _configure_github_proxy from .middleware import ( SanitizeToolInputsMiddleware, + SlackAssistantStatusMiddleware, ToolErrorMiddleware, check_message_queue_before_model, ensure_no_empty_msg, @@ -371,6 +372,7 @@ async def get_agent(config: RunnableConfig) -> Pregel: ModelCallLimitMiddleware(run_limit=MODEL_CALL_RECURSION_LIMIT, exit_behavior="end"), ToolErrorMiddleware(), check_message_queue_before_model, + SlackAssistantStatusMiddleware(), ensure_no_empty_msg, notify_step_limit_reached, ], diff --git a/agent/utils/slack.py b/agent/utils/slack.py index 221ab310..78099f1c 100644 --- a/agent/utils/slack.py +++ b/agent/utils/slack.py @@ -26,6 +26,30 @@ GITHUB_PR_URL_RE = re.compile(r"https?://(?:www\.)?github\.com/[^\s<>|]+/[^\s<>| URL_RE = re.compile(r"https?://[^\s<>|]+") +def _is_slack_assistants_api_enabled() -> bool: + """Whether the Slack Assistants API integration is enabled. + + Read at call time so tests and runtime can toggle via env without reimports. + """ + return os.environ.get("SLACK_ASSISTANTS_API_ENABLED", "").lower() in {"1", "true", "yes"} + + +DEFAULT_ASSISTANT_STATUS = "is thinking…" + +# Curated rotating loading strings shown by Slack while the indicator is active. +# Capped at 10 by Slack's API. +DEFAULT_LOADING_MESSAGES: tuple[str, ...] = ( + "reading the repo…", + "tracing call sites…", + "thinking through edge cases…", + "running commands…", + "drafting changes…", + "double-checking the diff…", + "writing tests…", + "tidying up…", +) + + @dataclass(frozen=True) class GitHubPrRef: owner: str @@ -255,6 +279,58 @@ def format_slack_messages_for_prompt( return "\n".join(lines) +async def set_slack_assistant_status( + channel_id: str, + thread_ts: str, + status: str = DEFAULT_ASSISTANT_STATUS, + loading_messages: list[str] | tuple[str, ...] | None = None, +) -> bool: + """Set the assistant typing/status indicator on a Slack thread. + + Wraps Slack's `assistants.threads.setStatus` API. The `chat:write` scope + on the bot token is sufficient. Status auto-clears when the bot posts to + the thread, and Slack itself expires it after ~2 minutes — callers that + want it visible across longer runs must refresh it periodically. + + `loading_messages` is an optional list (max 10) of strings Slack rotates + through while the indicator is visible. + + No-op (returning False) when the assistants feature flag is disabled, + the bot token is missing, or the channel/thread is not provided. + Failures are logged but never raised — the indicator is a UX nicety, + not a correctness requirement. + """ + if not _is_slack_assistants_api_enabled(): + return False + if not SLACK_BOT_TOKEN or not channel_id or not thread_ts: + return False + + payload: dict[str, Any] = { + "channel_id": channel_id, + "thread_ts": thread_ts, + "status": status, + } + if loading_messages: + payload["loading_messages"] = list(loading_messages)[:10] + + async with httpx.AsyncClient() as http_client: + try: + response = await http_client.post( + f"{SLACK_API_BASE_URL}/assistants.threads.setStatus", + headers=_slack_headers(), + json=payload, + ) + response.raise_for_status() + data = response.json() + if not data.get("ok"): + logger.warning("Slack assistants.threads.setStatus failed: %s", data.get("error")) + return False + return True + except httpx.HTTPError: + logger.exception("Slack assistants.threads.setStatus request failed") + return False + + async def post_slack_thread_reply(channel_id: str, thread_ts: str, text: str) -> bool: """Post a reply in a Slack thread.""" if not SLACK_BOT_TOKEN: diff --git a/agent/webapp.py b/agent/webapp.py index 38811595..b312bd4c 100644 --- a/agent/webapp.py +++ b/agent/webapp.py @@ -64,6 +64,7 @@ from .utils.slack import ( post_slack_trace_reply, resolve_slack_links_in_context, select_slack_context_messages, + set_slack_assistant_status, strip_bot_mention, verify_slack_signature, ) @@ -766,6 +767,8 @@ async def process_slack_mention(event_data: dict[str, Any], repo_config: dict[st channel_id, ) + await set_slack_assistant_status(channel_id, thread_ts) + thread_id = generate_thread_id_from_slack_thread(channel_id, thread_ts) user_email = None @@ -906,6 +909,7 @@ async def process_slack_mention(event_data: dict[str, Any], repo_config: dict[st ) if is_first_mention: await post_slack_trace_reply(channel_id, thread_ts, thread_id) + await set_slack_assistant_status(channel_id, thread_ts) else: logger.info( "Skipping Slack trace reply for thread %s — agent will reply when run completes", @@ -916,6 +920,7 @@ async def process_slack_mention(event_data: dict[str, Any], repo_config: dict[st async def process_slack_pr_review_request( pr_ref: GitHubPrRef, channel_id: str, thread_ts: str ) -> None: + await set_slack_assistant_status(channel_id, thread_ts) result = await trigger_pr_review_from_ref( pr_ref, source="slack", @@ -928,6 +933,7 @@ async def process_slack_pr_review_request( await post_slack_trace_reply( channel_id, thread_ts, thread_id, message="Taking a look..." ) + await set_slack_assistant_status(channel_id, thread_ts) return await post_slack_thread_reply( @@ -1459,6 +1465,8 @@ async def trigger_pr_review_from_ref( base_sha=base_sha, head_sha=head_sha, branch_name=branch_name, + slack_channel_id=slack_channel_id, + slack_thread_ts=slack_thread_ts, ) thread_active = await is_thread_active(thread_id) @@ -1491,6 +1499,8 @@ def _build_reviewer_configurable( branch_name: str, re_review: bool = False, last_reviewed_sha: str = "", + slack_channel_id: str = "", + slack_thread_ts: str = "", ) -> dict[str, Any]: """Assemble the runnable-config ``configurable`` dict for a reviewer run.""" configurable: dict[str, Any] = { @@ -1509,6 +1519,11 @@ def _build_reviewer_configurable( configurable["branch_name"] = branch_name if last_reviewed_sha: configurable["last_reviewed_sha"] = last_reviewed_sha + if slack_channel_id and slack_thread_ts: + configurable["slack_thread"] = { + "channel_id": slack_channel_id, + "thread_ts": slack_thread_ts, + } return configurable diff --git a/tests/test_github_issue_webhook.py b/tests/test_github_issue_webhook.py index ad528f3a..d98591d9 100644 --- a/tests/test_github_issue_webhook.py +++ b/tests/test_github_issue_webhook.py @@ -613,8 +613,13 @@ def test_process_slack_pr_review_request_posts_trace_reply(monkeypatch) -> None: "message": message, } + async def fake_set_slack_assistant_status(channel_id: str, thread_ts: str) -> bool: + captured.setdefault("status_calls", []).append((channel_id, thread_ts)) + return True + monkeypatch.setattr(webapp, "trigger_pr_review_from_ref", fake_trigger_pr_review_from_ref) monkeypatch.setattr(webapp, "post_slack_trace_reply", fake_post_slack_trace_reply) + monkeypatch.setattr(webapp, "set_slack_assistant_status", fake_set_slack_assistant_status) asyncio.run( webapp.process_slack_pr_review_request( @@ -638,6 +643,10 @@ def test_process_slack_pr_review_request_posts_trace_reply(monkeypatch) -> None: "thread_id": "reviewer-thread-id", "message": "Taking a look...", } + assert captured["status_calls"] == [ + ("C123", "1700000000.000100"), + ("C123", "1700000000.000100"), + ] def test_process_github_pr_review_request_creates_reviewer_run(monkeypatch) -> None: @@ -783,6 +792,8 @@ def test_trigger_pr_review_from_ref_creates_reviewer_run(monkeypatch) -> None: url="https://github.com/langchain-ai/open-swe/pull/1244", ), source="slack", + slack_channel_id="C123", + slack_thread_ts="1700000000.000100", ) ) @@ -803,6 +814,10 @@ def test_trigger_pr_review_from_ref_creates_reviewer_run(monkeypatch) -> None: assert config["repo"] == {"owner": "langchain-ai", "name": "open-swe"} assert config["pr_number"] == 1244 assert config["review_requested"] is True + assert config["slack_thread"] == { + "channel_id": "C123", + "thread_ts": "1700000000.000100", + } def test_trigger_pr_review_from_ref_respects_reviewer_allowlist(monkeypatch) -> None: diff --git a/tests/test_refresh_slack_status_middleware.py b/tests/test_refresh_slack_status_middleware.py new file mode 100644 index 00000000..fe04940b --- /dev/null +++ b/tests/test_refresh_slack_status_middleware.py @@ -0,0 +1,221 @@ +import asyncio +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from langchain.agents.middleware.types import ModelRequest, ModelResponse +from langchain_core.messages import AIMessage +from langgraph.prebuilt.tool_node import ToolCallRequest + +from agent.middleware.refresh_slack_status import ( + SlackAssistantStatusMiddleware, + _status_from_recent_tool_calls, +) +from agent.utils.slack import DEFAULT_ASSISTANT_STATUS, DEFAULT_LOADING_MESSAGES + + +class TestSlackAssistantStatusMiddleware: + def _runtime(self) -> MagicMock: + return MagicMock() + + def _config(self) -> dict: + return {"configurable": {"slack_thread": {"channel_id": "C1", "thread_ts": "1.0"}}} + + @pytest.mark.asyncio + async def test_before_agent_sets_status_when_slack_thread_present(self) -> None: + middleware = SlackAssistantStatusMiddleware() + with ( + patch("agent.middleware.refresh_slack_status.get_config", return_value=self._config()), + patch( + "agent.middleware.refresh_slack_status.set_slack_assistant_status", + new_callable=AsyncMock, + ) as mock_set, + ): + result = await middleware.abefore_agent({"messages": []}, self._runtime()) + + assert result is None + mock_set.assert_awaited_once() + args = mock_set.await_args.args + kwargs = mock_set.await_args.kwargs + assert args == ("C1", "1.0") + assert kwargs["status"] == DEFAULT_ASSISTANT_STATUS + assert kwargs["loading_messages"] == list(DEFAULT_LOADING_MESSAGES) + + @pytest.mark.asyncio + async def test_after_agent_clears_status_when_slack_thread_present(self) -> None: + middleware = SlackAssistantStatusMiddleware() + with ( + patch("agent.middleware.refresh_slack_status.get_config", return_value=self._config()), + patch( + "agent.middleware.refresh_slack_status.set_slack_assistant_status", + new_callable=AsyncMock, + ) as mock_set, + ): + result = await middleware.aafter_agent({"messages": []}, self._runtime()) + + assert result is None + assert mock_set.await_args.kwargs["status"] == "" + assert mock_set.await_args.kwargs["loading_messages"] is None + + @pytest.mark.asyncio + async def test_model_call_uses_contextual_status_from_last_tool_call(self) -> None: + middleware = SlackAssistantStatusMiddleware() + ai = AIMessage( + content="", + tool_calls=[{"name": "grep", "args": {"pattern": "foo"}, "id": "tc1"}], + ) + request = ModelRequest( + model=MagicMock(), + messages=[], + state={"messages": [ai]}, + runtime=self._runtime(), + ) + + async def handler(_request: ModelRequest) -> ModelResponse: + return ModelResponse(result=[AIMessage(content="done")]) + + with ( + patch("agent.middleware.refresh_slack_status.get_config", return_value=self._config()), + patch( + "agent.middleware.refresh_slack_status.set_slack_assistant_status", + new_callable=AsyncMock, + ) as mock_set, + ): + response = await middleware.awrap_model_call(request, handler) + + assert response.result[0].content == "done" + assert mock_set.await_args_list[0].kwargs["status"] == "searching the codebase..." + + @pytest.mark.asyncio + async def test_tool_call_uses_tool_specific_status(self) -> None: + middleware = SlackAssistantStatusMiddleware() + request = ToolCallRequest( + tool_call={"name": "execute", "args": {}, "id": "tc1"}, + tool=MagicMock(), + state={}, + runtime=self._runtime(), + ) + + async def handler(_request: ToolCallRequest) -> MagicMock: + return MagicMock() + + with ( + patch("agent.middleware.refresh_slack_status.get_config", return_value=self._config()), + patch( + "agent.middleware.refresh_slack_status.set_slack_assistant_status", + new_callable=AsyncMock, + ) as mock_set, + ): + await middleware.awrap_tool_call(request, handler) + + assert mock_set.await_args_list[0].kwargs["status"] == "running commands..." + + @pytest.mark.asyncio + async def test_heartbeat_refreshes_while_call_is_running(self) -> None: + middleware = SlackAssistantStatusMiddleware( + heartbeat_interval_seconds=0.01, + max_heartbeat_seconds=10, + ) + refreshed = asyncio.Event() + set_count = 0 + real_sleep = asyncio.sleep + + async def fake_set_status(*_args: object, **_kwargs: object) -> bool: + nonlocal set_count + set_count += 1 + if set_count >= 2: + refreshed.set() + return True + + async def fake_sleep(_delay: float) -> None: + await real_sleep(0) + + request = ModelRequest( + model=MagicMock(), + messages=[], + state={"messages": []}, + runtime=self._runtime(), + ) + + async def handler(_request: ModelRequest) -> ModelResponse: + await refreshed.wait() + return ModelResponse(result=[AIMessage(content="done")]) + + with ( + patch("agent.middleware.refresh_slack_status.get_config", return_value=self._config()), + patch( + "agent.middleware.refresh_slack_status.set_slack_assistant_status", + side_effect=fake_set_status, + ), + patch("agent.middleware.refresh_slack_status.asyncio.sleep", side_effect=fake_sleep), + ): + response = await middleware.awrap_model_call(request, handler) + + assert response.result[0].content == "done" + assert set_count >= 2 + + @pytest.mark.asyncio + async def test_skips_when_slack_thread_missing(self) -> None: + middleware = SlackAssistantStatusMiddleware() + with ( + patch( + "agent.middleware.refresh_slack_status.get_config", + return_value={"configurable": {}}, + ), + patch( + "agent.middleware.refresh_slack_status.set_slack_assistant_status", + new_callable=AsyncMock, + ) as mock_set, + ): + result = await middleware.abefore_agent({"messages": []}, self._runtime()) + + assert result is None + mock_set.assert_not_called() + + @pytest.mark.asyncio + async def test_skips_when_channel_or_thread_blank(self) -> None: + middleware = SlackAssistantStatusMiddleware() + with ( + patch( + "agent.middleware.refresh_slack_status.get_config", + return_value={ + "configurable": {"slack_thread": {"channel_id": "", "thread_ts": "1.0"}} + }, + ), + patch( + "agent.middleware.refresh_slack_status.set_slack_assistant_status", + new_callable=AsyncMock, + ) as mock_set, + ): + result = await middleware.abefore_agent({"messages": []}, self._runtime()) + + assert result is None + mock_set.assert_not_called() + + def test_status_helper_falls_back_when_no_tool_calls(self) -> None: + assert _status_from_recent_tool_calls([]) == DEFAULT_ASSISTANT_STATUS + assert _status_from_recent_tool_calls([AIMessage(content="hi")]) == DEFAULT_ASSISTANT_STATUS + + def test_status_helper_falls_back_for_unknown_tool(self) -> None: + ai = AIMessage( + content="", + tool_calls=[{"name": "mystery_tool", "args": {}, "id": "tc1"}], + ) + assert _status_from_recent_tool_calls([ai]) == DEFAULT_ASSISTANT_STATUS + + @pytest.mark.asyncio + async def test_swallows_config_exceptions(self) -> None: + middleware = SlackAssistantStatusMiddleware() + with ( + patch( + "agent.middleware.refresh_slack_status.get_config", + side_effect=RuntimeError("boom"), + ), + patch( + "agent.middleware.refresh_slack_status.set_slack_assistant_status", + new_callable=AsyncMock, + ) as mock_set, + ): + result = await middleware.abefore_agent({"messages": []}, self._runtime()) + + assert result is None + mock_set.assert_not_called() diff --git a/tests/test_slack_assistants_status.py b/tests/test_slack_assistants_status.py new file mode 100644 index 00000000..6e6523c9 --- /dev/null +++ b/tests/test_slack_assistants_status.py @@ -0,0 +1,176 @@ +"""Tests for the Slack Assistants API integration (feature-flagged).""" + +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from agent.utils import slack as slack_utils + + +def _ok_response() -> MagicMock: + response = MagicMock() + response.json.return_value = {"ok": True} + response.raise_for_status.return_value = None + return response + + +def _err_response(error: str = "channel_not_found") -> MagicMock: + response = MagicMock() + response.json.return_value = {"ok": False, "error": error} + response.raise_for_status.return_value = None + return response + + +def _async_client_cm(post_response: MagicMock) -> AsyncMock: + client_cm = AsyncMock() + client_cm.__aenter__.return_value = client_cm + client_cm.post = AsyncMock(return_value=post_response) + return client_cm + + +@pytest.mark.asyncio +async def test_set_slack_assistant_status_noop_when_flag_disabled( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("SLACK_ASSISTANTS_API_ENABLED", raising=False) + monkeypatch.setattr(slack_utils, "SLACK_BOT_TOKEN", "xoxb-test") + + client_cm = _async_client_cm(_ok_response()) + with patch.object(slack_utils.httpx, "AsyncClient", return_value=client_cm): + ok = await slack_utils.set_slack_assistant_status("C1", "1.0", "thinking…") + + assert ok is False + client_cm.post.assert_not_called() + + +@pytest.mark.asyncio +async def test_set_slack_assistant_status_noop_when_no_token( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("SLACK_ASSISTANTS_API_ENABLED", "true") + monkeypatch.setattr(slack_utils, "SLACK_BOT_TOKEN", "") + + client_cm = _async_client_cm(_ok_response()) + with patch.object(slack_utils.httpx, "AsyncClient", return_value=client_cm): + ok = await slack_utils.set_slack_assistant_status("C1", "1.0", "thinking…") + + assert ok is False + client_cm.post.assert_not_called() + + +@pytest.mark.asyncio +async def test_set_slack_assistant_status_calls_correct_endpoint( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("SLACK_ASSISTANTS_API_ENABLED", "true") + monkeypatch.setattr(slack_utils, "SLACK_BOT_TOKEN", "xoxb-test") + + client_cm = _async_client_cm(_ok_response()) + with patch.object(slack_utils.httpx, "AsyncClient", return_value=client_cm): + ok = await slack_utils.set_slack_assistant_status("C1", "1.0", "thinking…") + + assert ok is True + client_cm.post.assert_awaited_once() + args, kwargs = client_cm.post.call_args + assert args[0].endswith("/assistants.threads.setStatus") + assert kwargs["json"] == { + "channel_id": "C1", + "thread_ts": "1.0", + "status": "thinking…", + } + + +@pytest.mark.asyncio +async def test_set_slack_assistant_status_passes_loading_messages( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("SLACK_ASSISTANTS_API_ENABLED", "true") + monkeypatch.setattr(slack_utils, "SLACK_BOT_TOKEN", "xoxb-test") + + client_cm = _async_client_cm(_ok_response()) + with patch.object(slack_utils.httpx, "AsyncClient", return_value=client_cm): + await slack_utils.set_slack_assistant_status( + "C1", "1.0", "thinking…", loading_messages=["a", "b", "c"] + ) + + _, kwargs = client_cm.post.call_args + assert kwargs["json"]["loading_messages"] == ["a", "b", "c"] + + +@pytest.mark.asyncio +async def test_set_slack_assistant_status_caps_loading_messages_at_10( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("SLACK_ASSISTANTS_API_ENABLED", "true") + monkeypatch.setattr(slack_utils, "SLACK_BOT_TOKEN", "xoxb-test") + + client_cm = _async_client_cm(_ok_response()) + with patch.object(slack_utils.httpx, "AsyncClient", return_value=client_cm): + await slack_utils.set_slack_assistant_status( + "C1", "1.0", loading_messages=[f"m{i}" for i in range(15)] + ) + + _, kwargs = client_cm.post.call_args + assert len(kwargs["json"]["loading_messages"]) == 10 + assert kwargs["json"]["loading_messages"][0] == "m0" + assert kwargs["json"]["loading_messages"][-1] == "m9" + + +@pytest.mark.asyncio +async def test_set_slack_assistant_status_omits_loading_messages_when_unset( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("SLACK_ASSISTANTS_API_ENABLED", "true") + monkeypatch.setattr(slack_utils, "SLACK_BOT_TOKEN", "xoxb-test") + + client_cm = _async_client_cm(_ok_response()) + with patch.object(slack_utils.httpx, "AsyncClient", return_value=client_cm): + await slack_utils.set_slack_assistant_status("C1", "1.0", "thinking…") + + _, kwargs = client_cm.post.call_args + assert "loading_messages" not in kwargs["json"] + + +@pytest.mark.asyncio +async def test_set_slack_assistant_status_returns_false_on_slack_error( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("SLACK_ASSISTANTS_API_ENABLED", "true") + monkeypatch.setattr(slack_utils, "SLACK_BOT_TOKEN", "xoxb-test") + + client_cm = _async_client_cm(_err_response("invalid_thread")) + with patch.object(slack_utils.httpx, "AsyncClient", return_value=client_cm): + ok = await slack_utils.set_slack_assistant_status("C1", "1.0", "thinking…") + + assert ok is False + + +@pytest.mark.asyncio +async def test_post_slack_thread_reply_does_not_call_set_status( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Slack auto-clears the indicator on post; no extra setStatus call needed.""" + monkeypatch.setenv("SLACK_ASSISTANTS_API_ENABLED", "true") + monkeypatch.setattr(slack_utils, "SLACK_BOT_TOKEN", "xoxb-test") + + client_cm = _async_client_cm(_ok_response()) + with patch.object(slack_utils.httpx, "AsyncClient", return_value=client_cm): + ok = await slack_utils.post_slack_thread_reply("C1", "1.0", "hello") + + assert ok is True + assert client_cm.post.await_count == 1 + assert client_cm.post.call_args.args[0].endswith("/chat.postMessage") + + +def test_is_slack_assistants_api_enabled_truthy_values( + monkeypatch: pytest.MonkeyPatch, +) -> None: + for truthy in ("1", "true", "TRUE", "yes", "Yes"): + monkeypatch.setenv("SLACK_ASSISTANTS_API_ENABLED", truthy) + assert slack_utils._is_slack_assistants_api_enabled() is True + + for falsy in ("0", "false", "no", "", "off"): + monkeypatch.setenv("SLACK_ASSISTANTS_API_ENABLED", falsy) + assert slack_utils._is_slack_assistants_api_enabled() is False