diff --git a/backend/chainlit/context.py b/backend/chainlit/context.py index e6cabfeab2..763d396c7a 100644 --- a/backend/chainlit/context.py +++ b/backend/chainlit/context.py @@ -7,6 +7,11 @@ from chainlit.session import ClientType, HTTPSession, WebsocketSession +# The event loop only keeps weak references to tasks, so a persistence task +# whose only reference was the create_task() call can be garbage collected +# before it reaches the data layer. Hold each one until it completes. +_persistence_tasks: set[asyncio.Task] = set() + if TYPE_CHECKING: from chainlit.emitter import BaseChainlitEmitter from chainlit.step import Step @@ -95,9 +100,11 @@ def init_http_context( if data_layer := get_data_layer(): if user_id := getattr(user, "id", None): - asyncio.create_task( + _task = asyncio.create_task( data_layer.update_thread(thread_id=thread_id, user_id=user_id) ) + _persistence_tasks.add(_task) + _task.add_done_callback(_persistence_tasks.discard) return context diff --git a/backend/chainlit/element.py b/backend/chainlit/element.py index 91e26adfed..dadf031d1d 100644 --- a/backend/chainlit/element.py +++ b/backend/chainlit/element.py @@ -24,6 +24,11 @@ from chainlit.data import get_data_layer from chainlit.logger import logger +# The event loop only keeps weak references to tasks, so a persistence task +# whose only reference was the create_task() call can be garbage collected +# before it reaches the data layer. Hold each one until it completes. +_persistence_tasks: set[asyncio.Task] = set() + mime_types = { "text": "text/plain", "tasklist": "application/json", @@ -212,7 +217,9 @@ async def _create(self, persist=True) -> bool: if (data_layer := get_data_layer()) and persist: try: - asyncio.create_task(data_layer.create_element(self)) + _task = asyncio.create_task(data_layer.create_element(self)) + _persistence_tasks.add(_task) + _task.add_done_callback(_persistence_tasks.discard) except Exception as e: logger.error(f"Failed to create element: {e!s}") if not self.url and (not self.chainlit_key or self.updatable): diff --git a/backend/chainlit/message.py b/backend/chainlit/message.py index 0700f630a7..24b3daa6c8 100644 --- a/backend/chainlit/message.py +++ b/backend/chainlit/message.py @@ -27,6 +27,11 @@ ) from chainlit.utils import utc_now +# The event loop only keeps weak references to tasks, so a persistence task +# whose only reference was the create_task() call can be garbage collected +# before it reaches the data layer. Hold each one until it completes. +_persistence_tasks: set[asyncio.Task] = set() + class MessageBase(ABC): id: str @@ -113,7 +118,9 @@ async def update( data_layer = get_data_layer() if data_layer: try: - asyncio.create_task(data_layer.update_step(step_dict)) + _task = asyncio.create_task(data_layer.update_step(step_dict)) + _persistence_tasks.add(_task) + _task.add_done_callback(_persistence_tasks.discard) except Exception as e: if self.fail_on_persist_error: raise e @@ -132,7 +139,9 @@ async def remove(self): data_layer = get_data_layer() if data_layer: try: - asyncio.create_task(data_layer.delete_step(step_dict["id"])) + _task = asyncio.create_task(data_layer.delete_step(step_dict["id"])) + _persistence_tasks.add(_task) + _task.add_done_callback(_persistence_tasks.discard) except Exception as e: if self.fail_on_persist_error: raise e @@ -147,7 +156,9 @@ async def _create(self): data_layer = get_data_layer() if data_layer and not self.persisted: try: - asyncio.create_task(data_layer.create_step(step_dict)) + _task = asyncio.create_task(data_layer.create_step(step_dict)) + _persistence_tasks.add(_task) + _task.add_done_callback(_persistence_tasks.discard) self.persisted = True except Exception as e: if self.fail_on_persist_error: diff --git a/backend/chainlit/step.py b/backend/chainlit/step.py index 1604bac8af..c2d7324757 100644 --- a/backend/chainlit/step.py +++ b/backend/chainlit/step.py @@ -18,6 +18,11 @@ from chainlit.types import FeedbackDict from chainlit.utils import utc_now +# The event loop only keeps weak references to tasks, so a persistence task +# whose only reference was the create_task() call can be garbage collected +# before it reaches the data layer. Hold each one until it completes. +_persistence_tasks: set[asyncio.Task] = set() + def check_add_step_in_cot(step: "Step"): is_message = step.type in [ @@ -342,7 +347,9 @@ async def update(self): if data_layer: try: - asyncio.create_task(data_layer.update_step(step_dict.copy())) + _task = asyncio.create_task(data_layer.update_step(step_dict.copy())) + _persistence_tasks.add(_task) + _task.add_done_callback(_persistence_tasks.discard) except Exception as e: if self.fail_on_persist_error: raise e @@ -367,7 +374,9 @@ async def remove(self): if data_layer: try: - asyncio.create_task(data_layer.delete_step(self.id)) + _task = asyncio.create_task(data_layer.delete_step(self.id)) + _persistence_tasks.add(_task) + _task.add_done_callback(_persistence_tasks.discard) except Exception as e: if self.fail_on_persist_error: raise e @@ -393,7 +402,9 @@ async def send(self): if data_layer: try: - asyncio.create_task(data_layer.create_step(step_dict.copy())) + _task = asyncio.create_task(data_layer.create_step(step_dict.copy())) + _persistence_tasks.add(_task) + _task.add_done_callback(_persistence_tasks.discard) self.persisted = True except Exception as e: if self.fail_on_persist_error: @@ -493,7 +504,9 @@ def __enter__(self): self.parent_id = parent_step.id local_steps.set(previous_steps + [self]) - asyncio.create_task(self.send()) + _task = asyncio.create_task(self.send()) + _persistence_tasks.add(_task) + _task.add_done_callback(_persistence_tasks.discard) return self def __exit__(self, exc_type, exc_val, exc_tb): @@ -508,4 +521,6 @@ def __exit__(self, exc_type, exc_val, exc_tb): current_steps.remove(self) local_steps.set(current_steps) - asyncio.create_task(self.update()) + _task = asyncio.create_task(self.update()) + _persistence_tasks.add(_task) + _task.add_done_callback(_persistence_tasks.discard)