diff --git a/slack_bolt/app/app.py b/slack_bolt/app/app.py index 0e1a397b0..efc95291d 100644 --- a/slack_bolt/app/app.py +++ b/slack_bolt/app/app.py @@ -28,6 +28,7 @@ IgnoringSelfEvents, CustomMiddleware, ) +from slack_bolt.middleware.message_listener_matches import MessageListenerMatches from slack_bolt.middleware.url_verification import UrlVerification from slack_bolt.oauth import OAuthFlow from slack_bolt.request import BoltRequest @@ -388,20 +389,11 @@ def message( middleware: Optional[List[Union[Callable, Middleware]]] = None, ): matchers = matchers if matchers else [] + middleware = middleware if middleware else [] def __call__(func): primary_matcher = builtin_matchers.event("message") - - def keyword_matcher(payload) -> bool: - text: Optional[str] = payload.get("event", {}).get("text", {}) - if text: - if isinstance(keyword, Pattern): - return keyword.match(text) # type: ignore - elif isinstance(keyword, str): - return keyword in text - return False - - matchers.insert(0, keyword_matcher) + middleware.append(MessageListenerMatches(keyword)) return self._register_listener( func, primary_matcher, matchers, middleware, True ) diff --git a/slack_bolt/app/async_app.py b/slack_bolt/app/async_app.py index 91d6eb702..2d0798143 100644 --- a/slack_bolt/app/async_app.py +++ b/slack_bolt/app/async_app.py @@ -33,6 +33,9 @@ AsyncMiddleware, AsyncCustomMiddleware, ) +from slack_bolt.middleware.async_message_listener_matches import ( + AsyncMessageListenerMatches, +) from slack_bolt.middleware.authorization.async_multi_teams_authorization import ( AsyncMultiTeamsAuthorization, ) @@ -418,21 +421,11 @@ def message( middleware: Optional[List[Union[Callable, AsyncMiddleware]]] = None, ): matchers = matchers if matchers else [] + middleware = middleware if middleware else [] def __call__(func): primary_matcher = builtin_matchers.event("message", True) - - async def keyword_matcher(payload) -> bool: - text: Optional[str] = payload.get("event", {}).get("text", {}) - if text: - if isinstance(keyword, Pattern): - return keyword.match(text) # type: ignore - elif isinstance(keyword, str): - return keyword in text - return False - - matchers.insert(0, keyword_matcher) - + middleware.append(AsyncMessageListenerMatches(keyword)) return self._register_listener( func, primary_matcher, matchers, middleware, True ) diff --git a/slack_bolt/context/base_context.py b/slack_bolt/context/base_context.py index ab7c2d596..77a806105 100644 --- a/slack_bolt/context/base_context.py +++ b/slack_bolt/context/base_context.py @@ -1,5 +1,5 @@ from logging import Logger -from typing import Optional +from typing import Optional, Tuple from slack_bolt.auth import AuthorizationResult from slack_sdk import WebClient @@ -41,3 +41,8 @@ def channel_id(self) -> Optional[str]: @property def response_url(self) -> Optional[str]: return self.get("response_url", None) + + @property + def matches(self) -> Optional[Tuple]: + """Returns all the matched parts in message listener's regexp""" + return self.get("matches", None) diff --git a/slack_bolt/middleware/async_message_listener_matches.py b/slack_bolt/middleware/async_message_listener_matches.py new file mode 100644 index 000000000..9d95a5cc1 --- /dev/null +++ b/slack_bolt/middleware/async_message_listener_matches.py @@ -0,0 +1,28 @@ +import re +from typing import Callable, Awaitable, Union, Pattern + +from slack_bolt.request.async_request import AsyncBoltRequest +from slack_bolt.response import BoltResponse +from .async_middleware import AsyncMiddleware + + +class AsyncMessageListenerMatches(AsyncMiddleware): + def __init__(self, keyword: Union[str, Pattern]): + self.keyword = keyword + + async def async_process( + self, + *, + req: AsyncBoltRequest, + resp: BoltResponse, + next: Callable[[], Awaitable[BoltResponse]], + ) -> BoltResponse: + text = req.payload.get("event", {}).get("text", "") + if text: + m = re.search(self.keyword, text) + if m is not None: + req.context["matches"] = m.groups() # tuple + return await next() + + # As the text doesn't match, skip running the listener + return resp diff --git a/slack_bolt/middleware/message_listener_matches.py b/slack_bolt/middleware/message_listener_matches.py new file mode 100644 index 000000000..a63e458e1 --- /dev/null +++ b/slack_bolt/middleware/message_listener_matches.py @@ -0,0 +1,24 @@ +import re +from typing import Callable, Pattern, Union + +from slack_bolt.request import BoltRequest +from slack_bolt.response import BoltResponse +from .middleware import Middleware + + +class MessageListenerMatches(Middleware): # type: ignore + def __init__(self, keyword: Union[str, Pattern]): + self.keyword = keyword + + def process( + self, *, req: BoltRequest, resp: BoltResponse, next: Callable[[], BoltResponse], + ) -> BoltResponse: + text = req.payload.get("event", {}).get("text", "") + if text: + m = re.search(self.keyword, text) + if m is not None: + req.context["matches"] = m.groups() # tuple + return next() + + # As the text doesn't match, skip running the listener + return resp diff --git a/tests/async_scenario_tests/test_message.py b/tests/async_scenario_tests/test_message.py index d1c971cd0..e04cf02e8 100644 --- a/tests/async_scenario_tests/test_message.py +++ b/tests/async_scenario_tests/test_message.py @@ -51,6 +51,10 @@ def build_request(self) -> AsyncBoltRequest: timestamp, body = str(int(time())), json.dumps(message_payload) return AsyncBoltRequest(body=body, headers=self.build_headers(timestamp, body)) + def build_request2(self) -> AsyncBoltRequest: + timestamp, body = str(int(time())), json.dumps(message_payload2) + return AsyncBoltRequest(body=body, headers=self.build_headers(timestamp, body)) + @pytest.mark.asyncio async def test_string_keyword(self): app = AsyncApp(client=self.web_client, signing_secret=self.signing_secret,) @@ -63,6 +67,32 @@ async def test_string_keyword(self): await asyncio.sleep(1) # wait a bit after auto ack() assert self.mock_received_requests["/chat.postMessage"] == 1 + @pytest.mark.asyncio + async def test_string_keyword_capturing(self): + app = AsyncApp(client=self.web_client, signing_secret=self.signing_secret,) + app.message("We've received ([0-9]+) messages from (.+)!")(verify_matches) + + request = self.build_request2() + response = await app.async_dispatch(request) + assert response.status == 200 + assert self.mock_received_requests["/auth.test"] == 1 + await asyncio.sleep(1) # wait a bit after auto ack() + assert self.mock_received_requests["/chat.postMessage"] == 1 + + @pytest.mark.asyncio + async def test_string_keyword_capturing2(self): + app = AsyncApp(client=self.web_client, signing_secret=self.signing_secret,) + app.message(re.compile("We've received ([0-9]+) messages from (.+)!"))( + verify_matches + ) + + request = self.build_request2() + response = await app.async_dispatch(request) + assert response.status == 200 + assert self.mock_received_requests["/auth.test"] == 1 + await asyncio.sleep(1) # wait a bit after auto ack() + assert self.mock_received_requests["/chat.postMessage"] == 1 + @pytest.mark.asyncio async def test_string_keyword_unmatched(self): app = AsyncApp(client=self.web_client, signing_secret=self.signing_secret,) @@ -134,3 +164,32 @@ async def test_regexp_keyword_unmatched(self): async def whats_up(payload, say): assert payload == message_payload await say("What's up?") + + +message_payload2 = { + "token": "verification_token", + "team_id": "T111", + "enterprise_id": "E111", + "api_app_id": "A111", + "event": { + "client_msg_id": "a8744611-0210-4f85-9f15-5faf7fb225c8", + "type": "message", + "text": "We've received 103 messages from you!", + "user": "W111", + "ts": "1596183880.004200", + "team": "T111", + "channel": "C111", + "event_ts": "1596183880.004200", + "channel_type": "channel", + }, + "type": "event_callback", + "event_id": "Ev111", + "event_time": 1596183880, + "authed_users": ["W111"], +} + + +async def verify_matches(context, say): + assert context["matches"] == ("103", "you") + assert context.matches == ("103", "you") + await say("Thanks!") diff --git a/tests/scenario_tests/test_message.py b/tests/scenario_tests/test_message.py index 4f315c165..463c2621b 100644 --- a/tests/scenario_tests/test_message.py +++ b/tests/scenario_tests/test_message.py @@ -45,6 +45,10 @@ def build_request(self) -> BoltRequest: timestamp, body = str(int(time.time())), json.dumps(message_payload) return BoltRequest(body=body, headers=self.build_headers(timestamp, body)) + def build_request2(self) -> BoltRequest: + timestamp, body = str(int(time.time())), json.dumps(message_payload2) + return BoltRequest(body=body, headers=self.build_headers(timestamp, body)) + def test_string_keyword(self): app = App(client=self.web_client, signing_secret=self.signing_secret,) app.message("Hello")(whats_up) @@ -56,6 +60,30 @@ def test_string_keyword(self): time.sleep(1) # wait a bit after auto ack() assert self.mock_received_requests["/chat.postMessage"] == 1 + def test_string_keyword_capturing(self): + app = App(client=self.web_client, signing_secret=self.signing_secret,) + app.message("We've received ([0-9]+) messages from (.+)!")(verify_matches) + + request = self.build_request2() + response = app.dispatch(request) + assert response.status == 200 + assert self.mock_received_requests["/auth.test"] == 1 + time.sleep(1) # wait a bit after auto ack() + assert self.mock_received_requests["/chat.postMessage"] == 1 + + def test_string_keyword_capturing2(self): + app = App(client=self.web_client, signing_secret=self.signing_secret,) + app.message(re.compile("We've received ([0-9]+) messages from (.+)!"))( + verify_matches + ) + + request = self.build_request2() + response = app.dispatch(request) + assert response.status == 200 + assert self.mock_received_requests["/auth.test"] == 1 + time.sleep(1) # wait a bit after auto ack() + assert self.mock_received_requests["/chat.postMessage"] == 1 + def test_string_keyword_unmatched(self): app = App(client=self.web_client, signing_secret=self.signing_secret,) app.message("HELLO")(whats_up) @@ -124,3 +152,32 @@ def test_regexp_keyword_unmatched(self): def whats_up(payload, say): assert payload == message_payload say("What's up?") + + +message_payload2 = { + "token": "verification_token", + "team_id": "T111", + "enterprise_id": "E111", + "api_app_id": "A111", + "event": { + "client_msg_id": "a8744611-0210-4f85-9f15-5faf7fb225c8", + "type": "message", + "text": "We've received 103 messages from you!", + "user": "W111", + "ts": "1596183880.004200", + "team": "T111", + "channel": "C111", + "event_ts": "1596183880.004200", + "channel_type": "channel", + }, + "type": "event_callback", + "event_id": "Ev111", + "event_time": 1596183880, + "authed_users": ["W111"], +} + + +def verify_matches(context, say): + assert context["matches"] == ("103", "you") + assert context.matches == ("103", "you") + say("Thanks!")