Skip to content

Commit 0c31db2

Browse files
committed
feat: show a spinner when the assistant is busy
1 parent 3cb85ce commit 0c31db2

4 files changed

Lines changed: 167 additions & 57 deletions

File tree

src/generative_ai_toolkit/ui/conversation_list/conversation_list.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,7 @@
1919
from collections.abc import Sequence
2020
from dataclasses import dataclass
2121
from pathlib import Path
22-
from typing import TYPE_CHECKING, Any, Protocol, Unpack
22+
from typing import TYPE_CHECKING, Any, Protocol, Unpack, runtime_checkable
2323

2424
import boto3
2525
import boto3.session
@@ -51,6 +51,7 @@ class ConversationPage:
5151
next_page_token: Any | None = None
5252

5353

54+
@runtime_checkable
5455
class ConversationList(Protocol):
5556

5657
@property

src/generative_ai_toolkit/ui/lib.py

Lines changed: 47 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,7 @@ class TraceSummary:
3838
trace_id: str
3939
span_id: str
4040
started_at: datetime
41+
ended_at: datetime | None
4142
duration_ms: int | None
4243
conversation_id: str
4344
auth_context: AuthContext = field(default_factory=lambda: {"principal_id": None})
@@ -70,6 +71,7 @@ def get_summaries_for_traces(traces: Sequence[Trace]):
7071
span_id=root_trace.span_id,
7172
duration_ms=root_trace.ended_at and root_trace.duration_ms,
7273
started_at=root_trace.started_at,
74+
ended_at=root_trace.ended_at,
7375
all_traces=traces_for_trace_id,
7476
agent_cycle_traces={
7577
trace.span_id: trace
@@ -475,7 +477,11 @@ def chat_messages_from_trace_summary(
475477
)
476478
if cycle_response:
477479
metadata = metadata.copy()
478-
metadata.pop("status", None)
480+
if not trace.ended_at:
481+
metadata["status"] = "pending"
482+
elif metadata.get("status") == "done":
483+
# Always fold open
484+
metadata.pop("status")
479485
chat_messages.append(
480486
gr.ChatMessage(
481487
role="assistant",
@@ -537,7 +543,11 @@ def chat_messages_from_trace_summary(
537543
agent_response = trace.attributes.get("ai.agent.cycle.response")
538544
if agent_response:
539545
metadata = get_metadata(trace)
540-
metadata.pop("status", None)
546+
if not trace.ended_at:
547+
metadata["status"] = "pending"
548+
elif metadata.get("status") == "done":
549+
# Always fold open
550+
metadata.pop("status")
541551
chat_messages.append(
542552
gr.ChatMessage(
543553
role="assistant",
@@ -553,20 +563,29 @@ def chat_messages_from_trace_summary(
553563
return chat_messages
554564

555565

566+
@dataclass
567+
class ChatMessages:
568+
conversation_id: str
569+
principal_id: str | None
570+
messages: Sequence[gr.ChatMessage]
571+
assistant_busy: bool
572+
573+
556574
def chat_messages_from_traces(
557575
traces: Iterable[Trace],
558576
show_traces: Literal["ALL", "CORE", "CONVERSATION_ONLY"] = "CORE",
559577
):
560578
traces = list(traces)
561579
if not traces:
562-
return None, None, []
580+
return ChatMessages("", None, [], False)
563581
summaries = get_summaries_for_traces(traces)
564582
conversations = {
565583
(s.conversation_id, s.auth_context["principal_id"]) for s in summaries
566584
}
567585
if len(conversations) > 1:
568586
raise ValueError("More than one conversation id found")
569-
conversation_id, auth_context = conversations.pop()
587+
conversation_id, principal_id = conversations.pop()
588+
assistant_busy = not bool(summaries and summaries[-1].ended_at)
570589
messages = [
571590
msg
572591
for summary in summaries
@@ -575,7 +594,7 @@ def chat_messages_from_traces(
575594
include_traces=show_traces,
576595
)
577596
]
578-
return conversation_id, auth_context, messages
597+
return ChatMessages(conversation_id, principal_id, messages, assistant_busy)
579598

580599

581600
def chat_messages_from_conversation_measurements(
@@ -679,3 +698,26 @@ def format_date(dt: datetime):
679698
) # "Today" / "Yesterday" / "Monday"
680699

681700
return f"{day_text} at {dt.strftime("%X")}"
701+
702+
703+
def find_nearest_folded_open_message(messages: Sequence[gr.ChatMessage]):
704+
search_from = 0
705+
message = messages[-1]
706+
while message:
707+
message_parent_id = message.metadata.get("parent_id")
708+
if message.metadata.get("status") != "done": # Folded open!
709+
return message.metadata.get("id")
710+
elif message_parent_id:
711+
offset, message = next(
712+
(
713+
enumerate(
714+
msg
715+
for msg in reversed(messages[: len(messages) - search_from])
716+
if msg.metadata.get("id") == message_parent_id
717+
)
718+
),
719+
(-1, None),
720+
)
721+
search_from += offset
722+
continue
723+
return

0 commit comments

Comments
 (0)