diff --git a/slack_bolt/middleware/assistant/assistant.py b/slack_bolt/middleware/assistant/assistant.py index 98347d012..6e36651f7 100644 --- a/slack_bolt/middleware/assistant/assistant.py +++ b/slack_bolt/middleware/assistant/assistant.py @@ -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, @@ -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] @@ -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) @@ -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, diff --git a/slack_bolt/middleware/assistant/async_assistant.py b/slack_bolt/middleware/assistant/async_assistant.py index 588de8b41..87e837bef 100644 --- a/slack_bolt/middleware/assistant/async_assistant.py +++ b/slack_bolt/middleware/assistant/async_assistant.py @@ -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, @@ -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] @@ -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) @@ -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, diff --git a/tests/scenario_tests/test_events_assistant.py b/tests/scenario_tests/test_events_assistant.py index c95296154..8f37af5ae 100644 --- a/tests/scenario_tests/test_events_assistant.py +++ b/tests/scenario_tests/test_events_assistant.py @@ -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() diff --git a/tests/scenario_tests_async/test_events_assistant.py b/tests/scenario_tests_async/test_events_assistant.py index 9e1176c74..a830346cd 100644 --- a/tests/scenario_tests_async/test_events_assistant.py +++ b/tests/scenario_tests_async/test_events_assistant.py @@ -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)