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
9 changes: 8 additions & 1 deletion backend/chainlit/context.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand Down
9 changes: 8 additions & 1 deletion backend/chainlit/element.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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):
Expand Down
17 changes: 14 additions & 3 deletions backend/chainlit/message.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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:
Expand Down
25 changes: 20 additions & 5 deletions backend/chainlit/step.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 [
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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:
Expand Down Expand Up @@ -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):
Expand All @@ -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)