diff --git a/slack_sdk/web/async_chat_stream.py b/slack_sdk/web/async_chat_stream.py index 7348b90bc..dd4c0b459 100644 --- a/slack_sdk/web/async_chat_stream.py +++ b/slack_sdk/web/async_chat_stream.py @@ -120,7 +120,7 @@ async def append( self._token = kwargs.pop("token") if markdown_text is not None: self._buffer += markdown_text - if len(self._buffer) >= self._buffer_size or chunks is not None: + if not self._stream_ts or len(self._buffer) >= self._buffer_size or chunks is not None: return await self._flush_buffer(chunks=chunks, **kwargs) details = { "buffer_length": len(self._buffer), @@ -179,14 +179,10 @@ async def stop( if markdown_text: self._buffer += markdown_text if not self._stream_ts: - response = await self._client.chat_startStream( - **self._stream_args, - token=self._token, - ) - if not response.get("ts"): + await self._flush_buffer(chunks=chunks, **kwargs) + if not self._stream_ts: raise e.SlackRequestError("Failed to stop stream: stream not started") - self._stream_ts = str(response["ts"]) - self._state = "in_progress" + chunks = None flushings: List[Union[Dict, Chunk]] = [] if len(self._buffer) != 0: flushings.append(MarkdownTextChunk(text=self._buffer)) diff --git a/slack_sdk/web/chat_stream.py b/slack_sdk/web/chat_stream.py index 683859490..bee126d24 100644 --- a/slack_sdk/web/chat_stream.py +++ b/slack_sdk/web/chat_stream.py @@ -110,7 +110,7 @@ def append( self._token = kwargs.pop("token") if markdown_text is not None: self._buffer += markdown_text - if len(self._buffer) >= self._buffer_size or chunks is not None: + if not self._stream_ts or len(self._buffer) >= self._buffer_size or chunks is not None: return self._flush_buffer(chunks=chunks, **kwargs) details = { "buffer_length": len(self._buffer), @@ -169,14 +169,10 @@ def stop( if markdown_text: self._buffer += markdown_text if not self._stream_ts: - response = self._client.chat_startStream( - **self._stream_args, - token=self._token, - ) - if not response.get("ts"): + self._flush_buffer(chunks=chunks, **kwargs) + if not self._stream_ts: raise e.SlackRequestError("Failed to stop stream: stream not started") - self._stream_ts = str(response["ts"]) - self._state = "in_progress" + chunks = None flushings: List[Union[Dict, Chunk]] = [] if len(self._buffer) != 0: flushings.append(MarkdownTextChunk(text=self._buffer)) diff --git a/tests/slack_sdk/web/test_chat_stream.py b/tests/slack_sdk/web/test_chat_stream.py index 0a11b9d53..cfb86573e 100644 --- a/tests/slack_sdk/web/test_chat_stream.py +++ b/tests/slack_sdk/web/test_chat_stream.py @@ -103,14 +103,15 @@ def test_streams_a_short_message(self): self.assertEqual(start_request.get("recipient_team_id"), "T0123456789") self.assertEqual(start_request.get("recipient_user_id"), "U0123456789") - stop_request = self.thread.server.chat_stream_requests.get("/chat.stopStream", {}) - self.assertEqual(stop_request.get("channel"), "C0123456789") - self.assertEqual(stop_request.get("ts"), "123.123") self.assertEqual( - json.dumps(stop_request.get("chunks")), + json.dumps(start_request.get("chunks")), '[{"text": "nice!", "type": "markdown_text"}]', ) + stop_request = self.thread.server.chat_stream_requests.get("/chat.stopStream", {}) + self.assertEqual(stop_request.get("channel"), "C0123456789") + self.assertEqual(stop_request.get("ts"), "123.123") + def test_streams_a_long_message(self): streamer = self.client.chat_stream( buffer_size=5, @@ -211,7 +212,7 @@ def test_streams_a_chunk_message(self): ) self.assertEqual(self.received_requests.get("/chat.startStream", 0), 1) - self.assertEqual(self.received_requests.get("/chat.appendStream", 0), 1) + self.assertEqual(self.received_requests.get("/chat.appendStream", 0), 2) self.assertEqual(self.received_requests.get("/chat.stopStream", 0), 1) if hasattr(self.thread.server, "chat_stream_requests"): @@ -220,20 +221,12 @@ def test_streams_a_chunk_message(self): self.assertEqual(start_request.get("thread_ts"), "123.000") self.assertEqual( json.dumps(start_request.get("chunks")), - '[{"text": "**this is buffered**", "type": "markdown_text"}, {"id": "001", "status": "pending", "title": "Counting...", "type": "task_update"}]', + '[{"text": "**this is ", "type": "markdown_text"}]', ) self.assertEqual(start_request.get("recipient_team_id"), "T0123456789") self.assertEqual(start_request.get("recipient_user_id"), "U0123456789") self.assertEqual(start_request.get("task_display_mode"), "timeline") - append_request = self.thread.server.chat_stream_requests.get("/chat.appendStream", {}) - self.assertEqual(append_request.get("channel"), "C0123456789") - self.assertEqual(append_request.get("ts"), "123.123") - self.assertEqual( - json.dumps(append_request.get("chunks")), - '[{"text": "**this is unbuffered**", "type": "markdown_text"}]', - ) - stop_request = self.thread.server.chat_stream_requests.get("/chat.stopStream", {}) self.assertEqual(stop_request.get("channel"), "C0123456789") self.assertEqual(stop_request.get("ts"), "123.123") diff --git a/tests/slack_sdk_async/web/test_async_chat_stream.py b/tests/slack_sdk_async/web/test_async_chat_stream.py index 2a4f5b931..911f12742 100644 --- a/tests/slack_sdk_async/web/test_async_chat_stream.py +++ b/tests/slack_sdk_async/web/test_async_chat_stream.py @@ -105,14 +105,15 @@ async def test_streams_a_short_message(self): self.assertEqual(start_request.get("recipient_team_id"), "T0123456789") self.assertEqual(start_request.get("recipient_user_id"), "U0123456789") - stop_request = self.thread.server.chat_stream_requests.get("/chat.stopStream", {}) - self.assertEqual(stop_request.get("channel"), "C0123456789") - self.assertEqual(stop_request.get("ts"), "123.123") self.assertEqual( - json.dumps(stop_request.get("chunks")), + json.dumps(start_request.get("chunks")), '[{"text": "nice!", "type": "markdown_text"}]', ) + stop_request = self.thread.server.chat_stream_requests.get("/chat.stopStream", {}) + self.assertEqual(stop_request.get("channel"), "C0123456789") + self.assertEqual(stop_request.get("ts"), "123.123") + @async_test async def test_streams_a_long_message(self): streamer = await self.client.chat_stream( @@ -214,7 +215,7 @@ async def test_streams_a_chunk_message(self): ) self.assertEqual(self.received_requests.get("/chat.startStream", 0), 1) - self.assertEqual(self.received_requests.get("/chat.appendStream", 0), 1) + self.assertEqual(self.received_requests.get("/chat.appendStream", 0), 2) self.assertEqual(self.received_requests.get("/chat.stopStream", 0), 1) if hasattr(self.thread.server, "chat_stream_requests"): @@ -223,19 +224,11 @@ async def test_streams_a_chunk_message(self): self.assertEqual(start_request.get("thread_ts"), "123.000") self.assertEqual( json.dumps(start_request.get("chunks")), - '[{"text": "**this is buffered**", "type": "markdown_text"}, {"id": "001", "status": "pending", "title": "Counting...", "type": "task_update"}]', + '[{"text": "**this is ", "type": "markdown_text"}]', ) self.assertEqual(start_request.get("recipient_team_id"), "T0123456789") self.assertEqual(start_request.get("recipient_user_id"), "U0123456789") - append_request = self.thread.server.chat_stream_requests.get("/chat.appendStream", {}) - self.assertEqual(append_request.get("channel"), "C0123456789") - self.assertEqual(append_request.get("ts"), "123.123") - self.assertEqual( - json.dumps(append_request.get("chunks")), - '[{"text": "**this is unbuffered**", "type": "markdown_text"}]', - ) - stop_request = self.thread.server.chat_stream_requests.get("/chat.stopStream", {}) self.assertEqual(stop_request.get("channel"), "C0123456789") self.assertEqual(stop_request.get("ts"), "123.123")