diff --git a/app/frontend/components/Chat.tsx b/app/frontend/components/Chat.tsx
index ba0db1b0..7f069a05 100644
--- a/app/frontend/components/Chat.tsx
+++ b/app/frontend/components/Chat.tsx
@@ -152,6 +152,8 @@ function ToolCallIndicator({
return ;
case "get_document":
return ;
+ case "execute_skill":
+ return ;
default:
return ;
}
@@ -165,6 +167,8 @@ function ToolCallIndicator({
return "Ask";
case "get_document":
return "Document";
+ case "execute_skill":
+ return "Skill";
case "analyze":
return "Analyze";
case "research":
@@ -176,6 +180,16 @@ function ToolCallIndicator({
const getDescription = () => {
switch (toolName) {
+ case "execute_skill": {
+ const skill = args.skill_name as string | undefined;
+ const request = args.request as string | undefined;
+ return (
+
+ {skill ? `${skill}: ` : ""}
+ {request ?? "Processing..."}
+
+ );
+ }
case "search": {
const query = args.query as string;
return {query};
diff --git a/haiku_rag_slim/haiku/rag/chat/app.py b/haiku_rag_slim/haiku/rag/chat/app.py
index 6346c864..e0cd6142 100644
--- a/haiku_rag_slim/haiku/rag/chat/app.py
+++ b/haiku_rag_slim/haiku/rag/chat/app.py
@@ -1,5 +1,6 @@
# pyright: reportPossiblyUnboundVariable=false
import asyncio
+import json
import uuid
from collections.abc import Iterable
from datetime import datetime
@@ -32,6 +33,7 @@ try:
RunAgentInput,
StateDeltaEvent,
TextMessageContentEvent,
+ ToolCallArgsEvent,
ToolCallEndEvent,
ToolCallStartEvent,
UserMessage,
@@ -218,6 +220,7 @@ class ChatApp(App):
message = None
accumulated_text = ""
+ tool_args_deltas: dict[str, str] = {}
try:
async for event in adapter.run_stream():
@@ -249,7 +252,18 @@ class ChatApp(App):
await chat_history.add_tool_call(
event.tool_call_id, event.tool_call_name
)
+ tool_args_deltas[event.tool_call_id] = ""
await chat_history.show_thinking("Executing tasks...")
+ elif event.type == EventType.TOOL_CALL_ARGS:
+ assert isinstance(event, ToolCallArgsEvent)
+ tool_args_deltas[event.tool_call_id] = (
+ tool_args_deltas.get(event.tool_call_id, "") + event.delta
+ )
+ try:
+ args = json.loads(tool_args_deltas[event.tool_call_id])
+ chat_history.update_tool_args(event.tool_call_id, args)
+ except json.JSONDecodeError:
+ pass
elif event.type == EventType.TOOL_CALL_END:
assert isinstance(event, ToolCallEndEvent)
chat_history.mark_tool_complete(event.tool_call_id)
diff --git a/haiku_rag_slim/haiku/rag/chat/widgets/chat_history.py b/haiku_rag_slim/haiku/rag/chat/widgets/chat_history.py
index 90eb7c85..fdd4f940 100644
--- a/haiku_rag_slim/haiku/rag/chat/widgets/chat_history.py
+++ b/haiku_rag_slim/haiku/rag/chat/widgets/chat_history.py
@@ -59,7 +59,12 @@ class ToolCallWidget(Static):
yield Static(desc, classes="tool-desc")
def _build_description(self) -> str:
- if self.tool_name == "search":
+ if self.tool_name == "execute_skill":
+ skill = self.args.get("skill_name", "")
+ request = self.args.get("request", "...")
+ prefix = f"{skill}: " if skill else ""
+ return f'{prefix}"{request}"'
+ elif self.tool_name == "search":
query = self.args.get("query", "...")
return f'"{query}"'
elif self.tool_name == "ask":
@@ -72,6 +77,10 @@ class ToolCallWidget(Static):
return str(self.args)
return ""
+ def update_args(self, args: dict[str, Any]) -> None:
+ self.args = args
+ self.refresh(recompose=True)
+
def mark_completed(self) -> None:
self._completed = True
self.refresh(recompose=True)
@@ -333,6 +342,12 @@ class ChatHistory(VerticalScroll):
self.scroll_end(animate=False)
return widget
+ def update_tool_args(self, tool_call_id: str, args: dict[str, Any]) -> None:
+ """Update the args of a tool call widget."""
+ widget = self._tool_widgets.get(tool_call_id)
+ if widget:
+ widget.update_args(args)
+
def mark_tool_complete(self, tool_call_id: str) -> None:
"""Mark a tool call as complete by its ID."""
widget = self._tool_widgets.get(tool_call_id)