Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 16 additions & 4 deletions slack_bolt/middleware/assistant/assistant.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@
from slack_bolt.listener import Listener
from slack_bolt.listener.thread_runner import ThreadListenerRunner
from slack_bolt.middleware import Middleware
from slack_bolt.middleware.custom_middleware import CustomMiddleware
from slack_bolt.logger.messages import error_unexpected_listener_middleware
from slack_bolt.listener_matcher import ListenerMatcher
from slack_bolt.request.payload_utils import (
is_assistant_thread_started_event,
Expand Down Expand Up @@ -260,7 +262,7 @@ def build_listener(
self,
listener_or_functions: Union[Listener, Callable, List[Callable]],
matchers: Optional[List[Union[ListenerMatcher, Callable[..., bool]]]] = None,
middleware: Optional[List[Middleware]] = None,
middleware: Optional[List[Union[Callable, Middleware]]] = None,
base_logger: Optional[Logger] = None,
) -> Listener:
if isinstance(listener_or_functions, Callable): # type: ignore[arg-type]
Expand All @@ -269,8 +271,18 @@ def build_listener(
if isinstance(listener_or_functions, Listener):
return listener_or_functions
elif isinstance(listener_or_functions, list):
middleware = middleware if middleware else []
middleware.insert(0, AttachingConversationKwargs(self.thread_context_store))
# Build a new list so the caller's list is left untouched,
# and wrap plain functions the same way App does
listener_middleware: List[Middleware] = [AttachingConversationKwargs(self.thread_context_store)]
for m in middleware or []:
if isinstance(m, Middleware):
listener_middleware.append(m)
elif isinstance(m, Callable): # type: ignore[arg-type]
listener_middleware.append(
CustomMiddleware(app_name=self.app_name, func=m, base_logger=base_logger or self.base_logger)
)
else:
raise BoltError(error_unexpected_listener_middleware(type(m)))
functions = listener_or_functions
ack_function = functions.pop(0)

Expand All @@ -290,7 +302,7 @@ def build_listener(
return CustomListener(
app_name=self.app_name,
matchers=listener_matchers,
middleware=middleware,
middleware=listener_middleware,
ack_function=ack_function,
lazy_functions=functions,
auto_acknowledgement=True,
Expand Down
20 changes: 16 additions & 4 deletions slack_bolt/middleware/assistant/async_assistant.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@
from slack_bolt.error import BoltError
from slack_bolt.listener.async_listener import AsyncListener, AsyncCustomListener
from slack_bolt.middleware.async_middleware import AsyncMiddleware
from slack_bolt.middleware.async_custom_middleware import AsyncCustomMiddleware
from slack_bolt.logger.messages import error_unexpected_listener_middleware
from slack_bolt.listener_matcher.async_listener_matcher import AsyncListenerMatcher
from slack_bolt.request.payload_utils import (
is_assistant_thread_started_event,
Expand Down Expand Up @@ -293,7 +295,7 @@ def build_listener(
self,
listener_or_functions: Union[AsyncListener, Callable, List[Callable]],
matchers: Optional[List[Union[AsyncListenerMatcher, Callable[..., Awaitable[bool]]]]] = None,
middleware: Optional[List[AsyncMiddleware]] = None,
middleware: Optional[List[Union[Callable, AsyncMiddleware]]] = None,
base_logger: Optional[Logger] = None,
) -> AsyncListener:
if isinstance(listener_or_functions, Callable): # type: ignore[arg-type]
Expand All @@ -302,8 +304,18 @@ def build_listener(
if isinstance(listener_or_functions, AsyncListener):
return listener_or_functions
elif isinstance(listener_or_functions, list):
middleware = middleware if middleware else []
middleware.insert(0, AsyncAttachingConversationKwargs(self.thread_context_store))
# Build a new list so the caller's list is left untouched,
# and wrap plain functions the same way AsyncApp does
listener_middleware: List[AsyncMiddleware] = [AsyncAttachingConversationKwargs(self.thread_context_store)]
for m in middleware or []:
if isinstance(m, AsyncMiddleware):
listener_middleware.append(m)
elif isinstance(m, Callable): # type: ignore[arg-type]
listener_middleware.append(
AsyncCustomMiddleware(app_name=self.app_name, func=m, base_logger=base_logger or self.base_logger)
)
else:
raise BoltError(error_unexpected_listener_middleware(type(m)))
functions = listener_or_functions
ack_function = functions.pop(0)

Expand All @@ -323,7 +335,7 @@ def build_listener(
return AsyncCustomListener(
app_name=self.app_name,
matchers=listener_matchers,
middleware=middleware,
middleware=listener_middleware,
ack_function=ack_function,
lazy_functions=functions,
auto_acknowledgement=True,
Expand Down
39 changes: 39 additions & 0 deletions tests/scenario_tests/test_events_assistant.py
Original file line number Diff line number Diff line change
Expand Up @@ -226,6 +226,45 @@ def handle_user_message():
assert listener_called.wait(timeout=0.1) is True
assert middleware_called.wait(timeout=0.1) is True

def test_assistant_with_function_listener_middleware(self):
app = App(client=self.web_client)
assistant = Assistant()
listener_called = Event()
middleware_called = Event()

def function_middleware(next):
middleware_called.set()
next()

user_middleware = [function_middleware]

@assistant.thread_started(middleware=user_middleware)
def start_thread():
listener_called.set()

@assistant.user_message(middleware=user_middleware)
def handle_user_message():
listener_called.set()

app.assistant(assistant)
# the caller's list must not be modified
assert user_middleware == [function_middleware]

request = BoltRequest(body=thread_started_event_body, mode="socket_mode")
response = app.dispatch(request)
assert response.status == 200
assert listener_called.wait(timeout=0.1) is True
assert middleware_called.wait(timeout=0.1) is True

listener_called.clear()
middleware_called.clear()

request = BoltRequest(body=user_message_event_body, mode="socket_mode")
response = app.dispatch(request)
assert response.status == 200
assert listener_called.wait(timeout=0.1) is True
assert middleware_called.wait(timeout=0.1) is True

def test_assistant_custom_middleware_can_short_circuit(self):
app = App(client=self.web_client)
assistant = Assistant()
Expand Down
40 changes: 40 additions & 0 deletions tests/scenario_tests_async/test_events_assistant.py
Original file line number Diff line number Diff line change
Expand Up @@ -233,6 +233,46 @@ async def start_thread(context: AsyncBoltContext):
assert response.status == 200
assert (await asyncio.wait_for(listener_called.wait(), timeout=0.1)) is True

@pytest.mark.asyncio
async def test_assistant_with_function_listener_middleware(self):
app = AsyncApp(client=self.web_client)
assistant = AsyncAssistant()
listener_called = asyncio.Event()
middleware_called = asyncio.Event()

async def function_middleware(next):
middleware_called.set()
await next()

user_middleware = [function_middleware]

@assistant.thread_started(middleware=user_middleware)
async def start_thread():
listener_called.set()

@assistant.user_message(middleware=user_middleware)
async def handle_user_message():
listener_called.set()

app.assistant(assistant)
# the caller's list must not be modified
assert user_middleware == [function_middleware]

request = AsyncBoltRequest(body=thread_started_event_body, mode="socket_mode")
response = await app.async_dispatch(request)
assert response.status == 200
await asyncio.wait_for(listener_called.wait(), timeout=0.1)
await asyncio.wait_for(middleware_called.wait(), timeout=0.1)

listener_called.clear()
middleware_called.clear()

request = AsyncBoltRequest(body=user_message_event_body, mode="socket_mode")
response = await app.async_dispatch(request)
assert response.status == 200
await asyncio.wait_for(listener_called.wait(), timeout=0.1)
await asyncio.wait_for(middleware_called.wait(), timeout=0.1)

@pytest.mark.asyncio
async def test_assistant_with_custom_listener_middleware(self):
app = AsyncApp(client=self.web_client)
Expand Down