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)