From 2566bf1a61df2e7b9ef83088e968b0f9a4c5531a Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Mon, 2 Mar 2026 14:48:04 +0200 Subject: [PATCH] Adapt to haiku.skills streaming sub-agent ag-ui events --- CHANGELOG.md | 4 + app/backend/main.py | 15 +++- app/docker-compose.dev.yml | 1 + examples/custom_agent_agui.py | 11 ++- haiku_rag_slim/haiku/rag/chat/app.py | 121 ++++++++++++++------------- haiku_rag_slim/pyproject.toml | 2 +- uv.lock | 8 +- 7 files changed, 90 insertions(+), 72 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index d282aec5..fa22d108 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,10 @@ # Changelog ## [Unreleased] +### Changed + +- **AG-UI skill streaming**: Tool calls within skills are now streamed as real-time AG-UI events to the frontend. Requires `haiku.skills>=0.6.0` + ### Fixed - **Search tool regression**: Removed LLM-facing `filter` parameter from search and list_documents tools. The SQL WHERE clause description confused LLMs, degrading QA accuracy. Document filtering is now handled programmatically via `base_filter` and `state.document_filter` diff --git a/app/backend/main.py b/app/backend/main.py index d655d78e..c48c535b 100644 --- a/app/backend/main.py +++ b/app/backend/main.py @@ -18,7 +18,11 @@ from haiku.rag.config import load_yaml_config from haiku.rag.config.models import AppConfig from haiku.rag.skills.rag import AGENT_PREAMBLE, create_skill from haiku.rag.utils import get_model -from haiku.skills import SkillDeps, SkillToolset +from haiku.skills import ( + SkillDeps, + SkillToolset, + run_agui_stream, +) from haiku.skills.prompts import build_system_prompt load_dotenv(find_dotenv(usecwd=True)) @@ -83,11 +87,14 @@ async def stream_chat(request: Request) -> Response: run_input = AGUIAdapter.build_run_input(body) adapter = AGUIAdapter(agent=agent, run_input=run_input, accept=accept) - event_stream = adapter.run_stream(deps=SkillDeps()) - sse_event_stream = adapter.encode_stream(event_stream) + + async def event_stream(): + async with run_agui_stream(toolset, adapter, deps=SkillDeps()) as stream: + async for chunk in adapter.encode_stream(stream): + yield chunk return StreamingResponse( - sse_event_stream, + event_stream(), media_type=accept, headers={ "Cache-Control": "no-cache", diff --git a/app/docker-compose.dev.yml b/app/docker-compose.dev.yml index c80de2cb..80aa36b3 100644 --- a/app/docker-compose.dev.yml +++ b/app/docker-compose.dev.yml @@ -33,6 +33,7 @@ services: - "3000:3000" environment: - BACKEND_URL=http://backend:8000 + - HOSTNAME=0.0.0.0 volumes: - ./frontend:/app - frontend_node_modules:/app/node_modules diff --git a/examples/custom_agent_agui.py b/examples/custom_agent_agui.py index 105ca04a..bc27df25 100644 --- a/examples/custom_agent_agui.py +++ b/examples/custom_agent_agui.py @@ -24,7 +24,7 @@ from starlette.responses import JSONResponse, Response, StreamingResponse from starlette.routing import Route from haiku.rag.skills.rag import create_skill -from haiku.skills.agent import SkillToolset +from haiku.skills.agent import SkillToolset, run_agui_stream from haiku.skills.prompts import build_system_prompt db_path = os.environ.get("DB_PATH") @@ -50,11 +50,14 @@ async def stream_chat(request: Request) -> Response: run_input = AGUIAdapter.build_run_input(body) adapter = AGUIAdapter(agent=agent, run_input=run_input, accept=accept) - event_stream = adapter.run_stream() - sse_event_stream = adapter.encode_stream(event_stream) + + async def event_stream(): + async with run_agui_stream(toolset, adapter) as stream: + async for chunk in adapter.encode_stream(stream): + yield chunk return StreamingResponse( - sse_event_stream, + event_stream(), media_type=accept, headers={ "Cache-Control": "no-cache", diff --git a/haiku_rag_slim/haiku/rag/chat/app.py b/haiku_rag_slim/haiku/rag/chat/app.py index 7d8412db..da77de7a 100644 --- a/haiku_rag_slim/haiku/rag/chat/app.py +++ b/haiku_rag_slim/haiku/rag/chat/app.py @@ -10,7 +10,10 @@ from typing import TYPE_CHECKING, Any from haiku.rag.client import HaikuRAG from haiku.rag.config import get_config from haiku.rag.skills.rag import AGENT_PREAMBLE, RAGState -from haiku.skills.agent import SkillToolset +from haiku.skills.agent import ( + SkillToolset, + run_agui_stream, +) from haiku.skills.models import Skill from haiku.skills.prompts import build_system_prompt @@ -29,7 +32,6 @@ try: import textual_image.widget # noqa: F401 - import early for renderer detection from ag_ui.core import ( AssistantMessage, - BaseEvent, EventType, RunAgentInput, StateDeltaEvent, @@ -226,64 +228,65 @@ class ChatApp(App): tool_args_deltas: dict[str, str] = {} try: - async for event in adapter.run_stream(): - if not isinstance(event, BaseEvent): - continue - if event.type == EventType.TEXT_MESSAGE_START: - chat_history.hide_thinking() - message = await chat_history.add_message("assistant") - accumulated_text = "" - elif event.type == EventType.TEXT_MESSAGE_CONTENT: - assert isinstance(event, TextMessageContentEvent) - accumulated_text += event.delta - if message: - message.update_content(accumulated_text) - chat_history.scroll_end(animate=False) - elif event.type == EventType.TEXT_MESSAGE_END: - self._messages.append( - AssistantMessage( - id=str(uuid.uuid4()), - role="assistant", - content=accumulated_text, + async with run_agui_stream(self._toolset, adapter) as stream: + async for event in stream: + if event.type == EventType.TEXT_MESSAGE_START: + chat_history.hide_thinking() + message = await chat_history.add_message("assistant") + accumulated_text = "" + elif event.type == EventType.TEXT_MESSAGE_CONTENT: + assert isinstance(event, TextMessageContentEvent) + accumulated_text += event.delta + if message: + message.update_content(accumulated_text) + chat_history.scroll_end(animate=False) + elif event.type == EventType.TEXT_MESSAGE_END: + self._messages.append( + AssistantMessage( + id=str(uuid.uuid4()), + role="assistant", + content=accumulated_text, + ) + ) + # Show citations from RAG state + await self._show_citations(chat_history) + elif event.type == EventType.TOOL_CALL_START: + assert isinstance(event, ToolCallStartEvent) + chat_history.hide_thinking() + 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) + elif event.type == EventType.STATE_DELTA: + assert isinstance(event, StateDeltaEvent) + patch = JsonPatch(event.delta) + self._state = patch.apply(self._state) + self._toolset.restore_state_snapshot(self._state) + elif event.type == EventType.STATE_SNAPSHOT: + self._state = getattr(event, "snapshot", self._state) + self._toolset.restore_state_snapshot(self._state) + elif event.type == EventType.RUN_FINISHED: + chat_history.hide_thinking() + elif event.type == EventType.RUN_ERROR: + chat_history.hide_thinking() + error_msg = getattr(event, "message", "Unknown error") + await chat_history.add_message( + "assistant", f"Error: {error_msg}" ) - ) - # Show citations from RAG state - await self._show_citations(chat_history) - elif event.type == EventType.TOOL_CALL_START: - assert isinstance(event, ToolCallStartEvent) - chat_history.hide_thinking() - 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) - elif event.type == EventType.STATE_DELTA: - assert isinstance(event, StateDeltaEvent) - patch = JsonPatch(event.delta) - self._state = patch.apply(self._state) - self._toolset.restore_state_snapshot(self._state) - elif event.type == EventType.STATE_SNAPSHOT: - self._state = getattr(event, "snapshot", self._state) - self._toolset.restore_state_snapshot(self._state) - elif event.type == EventType.RUN_FINISHED: - chat_history.hide_thinking() - elif event.type == EventType.RUN_ERROR: - chat_history.hide_thinking() - error_msg = getattr(event, "message", "Unknown error") - await chat_history.add_message("assistant", f"Error: {error_msg}") except asyncio.CancelledError: chat_history.hide_thinking() diff --git a/haiku_rag_slim/pyproject.toml b/haiku_rag_slim/pyproject.toml index e8f32484..5159ff05 100644 --- a/haiku_rag_slim/pyproject.toml +++ b/haiku_rag_slim/pyproject.toml @@ -24,7 +24,7 @@ classifiers = [ dependencies = [ "cachetools>=5.5.0", "docling-core==2.65.1", - "haiku.skills>=0.5.1", + "haiku.skills>=0.6.0", "httpx>=0.28.1", "jsonpatch>=1.33", "lancedb==0.29.2", diff --git a/uv.lock b/uv.lock index 3c97235d..aab2e4c2 100644 --- a/uv.lock +++ b/uv.lock @@ -1490,7 +1490,7 @@ requires-dist = [ { name = "cohere", marker = "extra == 'cohere'", specifier = ">=5.20.1" }, { name = "docling", marker = "extra == 'docling'", specifier = "==2.73.1" }, { name = "docling-core", specifier = "==2.65.1" }, - { name = "haiku-skills", specifier = ">=0.5.1" }, + { name = "haiku-skills", specifier = ">=0.6.0" }, { name = "httpx", specifier = ">=0.28.1" }, { name = "jsonpatch", specifier = ">=1.33" }, { name = "lancedb", specifier = "==0.29.2" }, @@ -1522,7 +1522,7 @@ provides-extras = ["docling", "voyageai", "mxbai", "cohere", "zeroentropy", "jin [[package]] name = "haiku-skills" -version = "0.5.1" +version = "0.6.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "pydantic" }, @@ -1530,9 +1530,9 @@ dependencies = [ { name = "pyyaml" }, { name = "skills-ref" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/0b/4f/82d171478efe6c008705d3c2622009e5e5e18e6f4f2e4b9ad4916920d5d8/haiku_skills-0.5.1.tar.gz", hash = "sha256:8e46e7deac377a5f8fe7e4e61f8a3eca6328e069cd8b5aa63c4929a3b4f1bb6c", size = 122423, upload-time = "2026-02-27T09:31:30.785Z" } +sdist = { url = "https://files.pythonhosted.org/packages/10/0f/97d95a6814cec171d97ca7eec6365468f6cbdbfc9fbb2d6346494035aa8a/haiku_skills-0.6.0.tar.gz", hash = "sha256:8352c9157260742b475315f92191fe0c9149aa0faa8c9d20af8c201f1bc71e87", size = 132772, upload-time = "2026-03-03T08:55:59.909Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/18/98/99b1e5b0a2513abb82e7d8cd657d9ad548031c927344ca666eb92f588738/haiku_skills-0.5.1-py3-none-any.whl", hash = "sha256:662684ef13448c7b9c3c8c8d1f46256fe0a76bd5832dfb08a8f8933fc610e4c9", size = 23134, upload-time = "2026-02-27T09:31:29.624Z" }, + { url = "https://files.pythonhosted.org/packages/4c/c4/f9f8892da06bdc2defb12191e765486dd8dd87155c63dc6ca045f75c8211/haiku_skills-0.6.0-py3-none-any.whl", hash = "sha256:adfbe2206eb238abb0dfe376ffede669860ffd45a2f4647090dfbd787842cf7f", size = 24631, upload-time = "2026-03-03T08:55:58.901Z" }, ] [[package]]