Skip to content
Merged
Next Next commit
feat: make set_suggested_prompts accessible by all DMs to app
  • Loading branch information
WilliamBergamin committed Jul 2, 2026
commit ba8727251fc40d29124300d65ab9ce3274961672
Original file line number Diff line number Diff line change
Expand Up @@ -7,13 +7,13 @@
class AsyncSetSuggestedPrompts:
client: AsyncWebClient
channel_id: str
thread_ts: str
thread_ts: Optional[str]

def __init__(
self,
client: AsyncWebClient,
channel_id: str,
thread_ts: str,
thread_ts: Optional[str] = None,
):
self.client = client
self.channel_id = channel_id
Expand All @@ -23,6 +23,7 @@ async def __call__(
self,
prompts: Sequence[Union[str, Dict[str, str]]],
title: Optional[str] = None,
thread_ts: Optional[str] = None,
) -> AsyncSlackResponse:
prompts_arg: List[Dict[str, str]] = []
for prompt in prompts:
Expand All @@ -33,7 +34,7 @@ async def __call__(

return await self.client.assistant_threads_setSuggestedPrompts(
channel_id=self.channel_id,
thread_ts=self.thread_ts,
thread_ts=thread_ts or self.thread_ts,
prompts=prompts_arg,
title=title,
)
Original file line number Diff line number Diff line change
Expand Up @@ -7,13 +7,13 @@
class SetSuggestedPrompts:
client: WebClient
channel_id: str
thread_ts: str
thread_ts: Optional[str]

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

👍🏻


def __init__(
self,
client: WebClient,
channel_id: str,
thread_ts: str,
thread_ts: Optional[str] = None,
):
self.client = client
self.channel_id = channel_id
Expand All @@ -23,6 +23,7 @@ def __call__(
self,
prompts: Sequence[Union[str, Dict[str, str]]],
title: Optional[str] = None,
thread_ts: Optional[str] = None,
) -> SlackResponse:
prompts_arg: List[Dict[str, str]] = []
for prompt in prompts:
Expand All @@ -33,7 +34,7 @@ def __call__(

return self.client.assistant_threads_setSuggestedPrompts(
channel_id=self.channel_id,
thread_ts=self.thread_ts,
thread_ts=thread_ts or self.thread_ts,
prompts=prompts_arg,
title=title,
)
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,10 @@
from slack_bolt.context.assistant.thread_context_store.async_store import AsyncAssistantThreadContextStore
from slack_bolt.context.say_stream.async_say_stream import AsyncSayStream
from slack_bolt.context.set_status.async_set_status import AsyncSetStatus
from slack_bolt.context.set_suggested_prompts.async_set_suggested_prompts import AsyncSetSuggestedPrompts
from slack_bolt.middleware.async_middleware import AsyncMiddleware
from slack_bolt.request.async_request import AsyncBoltRequest
from slack_bolt.request.payload_utils import is_assistant_event, to_event
from slack_bolt.request.payload_utils import is_assistant_event, to_event, is_im_message_event
from slack_bolt.response import BoltResponse


Expand Down Expand Up @@ -38,19 +39,26 @@ async def async_process(
req.context["get_thread_context"] = assistant.get_thread_context
req.context["save_thread_context"] = assistant.save_thread_context

# TODO: in the future we might want to introduce a "proper" extract_ts utility
thread_ts = req.context.thread_ts or event.get("ts")
if req.context.channel_id and thread_ts:
req.context["set_status"] = AsyncSetStatus(
client=req.context.client,
channel_id=req.context.channel_id,
thread_ts=thread_ts,
)
req.context["say_stream"] = AsyncSayStream(
client=req.context.client,
channel=req.context.channel_id,
recipient_team_id=req.context.team_id or req.context.enterprise_id,
recipient_user_id=req.context.user_id,
thread_ts=thread_ts,
)
if req.context.channel_id:
# TODO: in the future we might want to introduce a "proper" extract_ts utility
thread_ts = req.context.thread_ts or event.get("ts")
if is_im_message_event(event):
req.context["set_suggested_prompts"] = AsyncSetSuggestedPrompts(
client=req.context.client,
channel_id=req.context.channel_id,
thread_ts=thread_ts,
)
if thread_ts:
req.context["set_status"] = AsyncSetStatus(
client=req.context.client,
channel_id=req.context.channel_id,
thread_ts=thread_ts,
)
req.context["say_stream"] = AsyncSayStream(
client=req.context.client,
channel=req.context.channel_id,
recipient_team_id=req.context.team_id or req.context.enterprise_id,
recipient_user_id=req.context.user_id,
thread_ts=thread_ts,
)
return await next()
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,9 @@
from slack_bolt.context.assistant.thread_context_store.store import AssistantThreadContextStore
from slack_bolt.context.say_stream.say_stream import SayStream
from slack_bolt.context.set_status.set_status import SetStatus
from slack_bolt.context.set_suggested_prompts.set_suggested_prompts import SetSuggestedPrompts
from slack_bolt.middleware import Middleware
from slack_bolt.request.payload_utils import is_assistant_event, to_event
from slack_bolt.request.payload_utils import is_assistant_event, is_im_message_event, to_event
from slack_bolt.request.request import BoltRequest
from slack_bolt.response.response import BoltResponse

Expand All @@ -32,19 +33,26 @@ def process(self, *, req: BoltRequest, resp: BoltResponse, next: Callable[[], Bo
req.context["get_thread_context"] = assistant.get_thread_context
req.context["save_thread_context"] = assistant.save_thread_context

# TODO: in the future we might want to introduce a "proper" extract_ts utility
thread_ts = req.context.thread_ts or event.get("ts")
if req.context.channel_id and thread_ts:
req.context["set_status"] = SetStatus(
client=req.context.client,
channel_id=req.context.channel_id,
thread_ts=thread_ts,
)
req.context["say_stream"] = SayStream(
client=req.context.client,
channel=req.context.channel_id,
recipient_team_id=req.context.team_id or req.context.enterprise_id,
recipient_user_id=req.context.user_id,
thread_ts=thread_ts,
)
if req.context.channel_id:
# TODO: in the future we might want to introduce a "proper" extract_ts utility
thread_ts = req.context.thread_ts or event.get("ts")
if is_im_message_event(event):
req.context["set_suggested_prompts"] = SetSuggestedPrompts(
client=req.context.client,
channel_id=req.context.channel_id,
thread_ts=thread_ts,
)
if thread_ts:
req.context["set_status"] = SetStatus(
client=req.context.client,
channel_id=req.context.channel_id,
thread_ts=thread_ts,
)
req.context["say_stream"] = SayStream(
client=req.context.client,
channel=req.context.channel_id,
recipient_team_id=req.context.team_id or req.context.enterprise_id,
recipient_user_id=req.context.user_id,
thread_ts=thread_ts,
)
return next()
20 changes: 13 additions & 7 deletions slack_bolt/request/payload_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@ def to_event(body: Dict[str, Any]) -> Optional[Dict[str, Any]]:


def to_message(body: Dict[str, Any]) -> Optional[Dict[str, Any]]:
if is_event(body) and body["event"]["type"] == "message":
if is_message_event(body):
return to_event(body)
return None

Expand All @@ -31,6 +31,12 @@ def is_workflow_step_execute(body: Dict[str, Any]) -> bool:
return is_event(body) and body["event"]["type"] == "workflow_step_execute" and "workflow_step" in body["event"]


def is_message_event(body: Dict[str, Any]) -> bool:
if is_event(body):
return body["event"]["type"] == "message"
return False


def is_assistant_event(body: Dict[str, Any]) -> bool:
return is_event(body) and (
is_assistant_thread_started_event(body)
Expand All @@ -52,16 +58,16 @@ def is_assistant_thread_context_changed_event(body: Dict[str, Any]) -> bool:
return False


def is_message_event_in_assistant_thread(body: Dict[str, Any]) -> bool:
if is_event(body):
return body["event"]["type"] == "message" and body["event"].get("channel_type") == "im"
def is_im_message_event(body: Dict[str, Any]) -> bool:
if is_message_event(body):
return body["event"].get("channel_type") == "im"
return False


def is_user_message_event_in_assistant_thread(body: Dict[str, Any]) -> bool:
if is_event(body):
return (
is_message_event_in_assistant_thread(body)
is_im_message_event(body)
and body["event"].get("subtype") in (None, "file_share")
and body["event"].get("thread_ts") is not None
and body["event"].get("bot_id") is None
Expand All @@ -72,7 +78,7 @@ def is_user_message_event_in_assistant_thread(body: Dict[str, Any]) -> bool:
def is_bot_message_event_in_assistant_thread(body: Dict[str, Any]) -> bool:
if is_event(body):
return (
is_message_event_in_assistant_thread(body)
is_im_message_event(body)
and body["event"].get("subtype") is None
and body["event"].get("thread_ts") is not None
and body["event"].get("bot_id") is not None
Expand All @@ -84,7 +90,7 @@ def is_other_message_sub_event_in_assistant_thread(body: Dict[str, Any]) -> bool
# message_changed, message_deleted etc.
if is_event(body):
return (
is_message_event_in_assistant_thread(body)
is_im_message_event(body)
and not is_user_message_event_in_assistant_thread(body)
and (
_is_other_message_sub_event(body["event"].get("message"))
Expand Down
10 changes: 10 additions & 0 deletions tests/slack_bolt/context/test_set_suggested_prompts.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,16 @@ def test_set_suggested_prompts_objects(self):
)
assert response.status_code == 200

def test_set_suggested_prompts_without_thread_ts(self):
set_suggested_prompts = SetSuggestedPrompts(client=self.web_client, channel_id="C111")
response: SlackResponse = set_suggested_prompts(prompts=["One", "Two"])
assert response.status_code == 200

def test_set_suggested_prompts_thread_ts_override(self):
set_suggested_prompts = SetSuggestedPrompts(client=self.web_client, channel_id="C111")
response: SlackResponse = set_suggested_prompts(prompts=["One", "Two"], thread_ts="123.123")
assert response.status_code == 200

def test_set_suggested_prompts_invalid(self):
set_suggested_prompts = SetSuggestedPrompts(client=self.web_client, channel_id="C111", thread_ts="123.123")
with pytest.raises(TypeError):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,18 @@ async def test_set_suggested_prompts_objects(self):
)
assert response.status_code == 200

@pytest.mark.asyncio
async def test_set_suggested_prompts_without_thread_ts(self):
set_suggested_prompts = AsyncSetSuggestedPrompts(client=self.web_client, channel_id="C111")
response: AsyncSlackResponse = await set_suggested_prompts(prompts=["One", "Two"])
assert response.status_code == 200

@pytest.mark.asyncio
async def test_set_suggested_prompts_thread_ts_override(self):
set_suggested_prompts = AsyncSetSuggestedPrompts(client=self.web_client, channel_id="C111")
response: AsyncSlackResponse = await set_suggested_prompts(prompts=["One", "Two"], thread_ts="123.123")
assert response.status_code == 200

@pytest.mark.asyncio
async def test_set_suggested_prompts_invalid(self):
set_suggested_prompts = AsyncSetSuggestedPrompts(client=self.web_client, channel_id="C111", thread_ts="123.123")
Expand Down