haiku.rag/haiku_rag_slim/haiku/rag/chat/app.py

341 lines
12 KiB
Python

# pyright: reportPossiblyUnboundVariable=false
import asyncio
import uuid
from collections.abc import AsyncIterable
from datetime import datetime
from pathlib import Path
from typing import TYPE_CHECKING, Any
from ag_ui.core import EventType
from pydantic_ai import (
Agent,
AgentStreamEvent,
FunctionToolCallEvent,
FunctionToolResultEvent,
RunContext,
)
from pydantic_ai.messages import ModelMessage
from haiku.rag.agents.chat.agent import create_chat_agent
from haiku.rag.agents.chat.state import (
AGUI_STATE_KEY,
ChatDeps,
ChatSessionState,
CitationInfo,
)
from haiku.rag.client import HaikuRAG
from haiku.rag.config import get_config
if TYPE_CHECKING:
from textual.app import ComposeResult
try:
import logfire
logfire.configure(send_to_logfire="if-token-present", console=False)
except ImportError:
pass
try:
import textual_image.widget # noqa: F401 - import early for renderer detection
from textual.app import App
from textual.binding import Binding
from textual.widgets import Footer, Header, Input
from textual.worker import Worker
from haiku.rag.chat.widgets.chat_history import ChatHistory, CitationWidget
TEXTUAL_AVAILABLE = True
except ImportError:
TEXTUAL_AVAILABLE = False
App = object # type: ignore
class ChatApp(App):
"""Textual TUI for conversational RAG."""
TITLE = "haiku.rag Chat"
CSS = """
Screen {
layout: grid;
grid-size: 1 2;
grid-rows: 1fr auto;
background: $surface;
}
#chat-history {
height: 100%;
}
Header {
background: $primary;
}
Footer {
background: $surface-darken-1;
}
"""
BINDINGS = [
Binding("ctrl+l", "clear_chat", "Clear", show=True),
Binding("ctrl+g", "show_visual", "Visual", show=True),
Binding("ctrl+i", "show_info", "Info", show=True),
Binding("escape", "focus_input", "Focus Input", show=False),
]
def __init__(
self,
db_path: Path,
read_only: bool = False,
before: datetime | None = None,
background_context: str | None = None,
) -> None:
super().__init__()
self.db_path = db_path
self.read_only = read_only
self.before = before
self.background_context = background_context
self.client: HaikuRAG | None = None
self.config = get_config()
self.agent: Agent[ChatDeps, str] | None = None
self.session_state: ChatSessionState | None = None
self._is_processing = False
self._tool_call_widgets: dict[str, Any] = {}
self._last_citations: list[CitationInfo] = []
self._selected_citation_idx: int | None = None
self._current_worker: Worker[None] | None = None
self._message_history: list[ModelMessage] = []
def compose(self) -> "ComposeResult":
"""Compose the UI layout."""
yield Header()
yield ChatHistory(id="chat-history")
yield Input(placeholder="Ask a question...", id="chat-input")
yield Footer()
async def on_mount(self) -> None:
"""Initialize the app when mounted."""
self.client = HaikuRAG(
db_path=self.db_path,
config=self.config,
read_only=self.read_only,
before=self.before,
)
await self.client.__aenter__()
# Create agent and session state
self.agent = create_chat_agent(self.config)
self.session_state = ChatSessionState(
session_id=str(uuid.uuid4()),
background_context=self.background_context,
)
# Focus the input field
self.query_one(Input).focus()
async def on_unmount(self) -> None:
"""Clean up when unmounting."""
if self.client:
await self.client.__aexit__(None, None, None)
async def _handle_stream_event(self, event: AgentStreamEvent) -> None:
"""Handle streaming events from the agent."""
chat_history = self.query_one(ChatHistory)
if isinstance(event, FunctionToolCallEvent):
tool_name = event.part.tool_name
tool_call_id = event.part.tool_call_id or str(uuid.uuid4())
args = event.part.args_as_dict()
widget = await chat_history.add_tool_call(tool_name, args)
self._tool_call_widgets[tool_call_id] = widget
elif isinstance(event, FunctionToolResultEvent):
tool_call_id = event.tool_call_id
if tool_call_id and tool_call_id in self._tool_call_widgets:
widget = self._tool_call_widgets[tool_call_id]
chat_history.mark_tool_complete(widget)
# Extract citations from StateSnapshotEvent in tool metadata
result = getattr(event, "result", None)
metadata = getattr(result, "metadata", None) if result else None
if metadata:
for meta_event in metadata:
if (
hasattr(meta_event, "type")
and meta_event.type == EventType.STATE_SNAPSHOT
):
snapshot = getattr(meta_event, "snapshot", {})
chat_state = snapshot.get(AGUI_STATE_KEY, snapshot)
self._last_citations = [
CitationInfo(**c) for c in chat_state["citations"]
]
async def _event_stream_handler(
self,
_ctx: RunContext[ChatDeps],
event_stream: AsyncIterable[AgentStreamEvent],
) -> None:
"""Handle streaming events from the agent."""
async for event in event_stream:
await self._handle_stream_event(event)
# Yield to event loop to keep UI responsive
await asyncio.sleep(0)
async def on_input_submitted(self, event: Input.Submitted) -> None:
"""Handle user input submission."""
user_message = event.value.strip()
if not user_message or self._is_processing:
return
if not self.client or not self.agent:
return
# Clear the input
event.input.clear()
# Add user message to history
chat_history = self.query_one(ChatHistory)
await chat_history.add_message("user", user_message)
# Clear for new query
self._tool_call_widgets.clear()
self._last_citations.clear()
self._selected_citation_idx = None
# Run agent in a worker to keep UI responsive
self._is_processing = True
self.query_one(Input).disabled = True
self._current_worker = self.run_worker(
self._run_agent(user_message), exclusive=True
)
async def _run_agent(self, user_message: str) -> None:
"""Run the agent in a background worker."""
if not self.client or not self.agent:
return
chat_history = self.query_one(ChatHistory)
# Show thinking indicator
await chat_history.show_thinking()
try:
deps = ChatDeps(
client=self.client,
config=self.config,
session_state=self.session_state,
state_key=AGUI_STATE_KEY,
)
async with self.agent.run_stream(
user_message,
deps=deps,
message_history=self._message_history,
event_stream_handler=self._event_stream_handler,
) as stream:
# Hide thinking when we start getting content
chat_history.hide_thinking()
# Create assistant message for streaming
assistant_msg = await chat_history.add_message("assistant", "")
# Stream text updates
async for text in stream.stream_text():
assistant_msg.update_content(text)
chat_history.scroll_end(animate=False)
# Yield to event loop to keep UI responsive
await asyncio.sleep(0)
# Update message history with this conversation
self._message_history = stream.all_messages()
# Add citations captured from tool metadata
if self._last_citations:
await chat_history.add_citations(self._last_citations)
except asyncio.CancelledError:
chat_history.hide_thinking()
await chat_history.add_message("assistant", "*Cancelled*")
except Exception as e:
chat_history.hide_thinking()
await chat_history.add_message("assistant", f"Error: {e}")
finally:
self._is_processing = False
self._current_worker = None
chat_input = self.query_one(Input)
chat_input.disabled = False
chat_input.focus()
async def action_clear_chat(self) -> None:
"""Clear the chat history and reset session."""
chat_history = self.query_one(ChatHistory)
await chat_history.clear_messages()
self._last_citations.clear()
self._selected_citation_idx = None
self._message_history.clear()
# Reset session state for fresh conversation (preserve background_context)
self.session_state = ChatSessionState(
session_id=str(uuid.uuid4()),
background_context=self.background_context,
)
def action_focus_input(self) -> None:
"""Focus the input field, or cancel if processing."""
if self._is_processing and self._current_worker:
self._current_worker.cancel()
self.query_one(Input).focus()
def _clear_citation_selection(self) -> None:
"""Clear citation selection."""
chat_history = self.query_one(ChatHistory)
for widget in chat_history.query(CitationWidget):
widget.remove_class("selected")
self._selected_citation_idx = None
def on_descendant_focus(self, _event: object) -> None:
"""Clear citation selection when input is focused."""
if isinstance(self.focused, Input):
self._clear_citation_selection()
async def action_show_visual(self) -> None:
"""Show visual grounding for the selected citation."""
if not self.client or not self._last_citations:
return
idx = (
self._selected_citation_idx
if self._selected_citation_idx is not None
else 0
)
citation = self._last_citations[idx]
chunk = await self.client.chunk_repository.get_by_id(citation.chunk_id)
if not chunk:
return
from haiku.rag.inspector.widgets.visual_modal import VisualGroundingModal
await self.push_screen(VisualGroundingModal(chunk=chunk, client=self.client))
async def action_show_info(self) -> None:
"""Show database info modal."""
if not self.client:
return
from haiku.rag.inspector.widgets.info_modal import InfoModal
await self.push_screen(InfoModal(self.client, self.db_path))
def on_citation_widget_selected(self, event: CitationWidget.Selected) -> None:
"""Handle citation selection."""
chat_history = self.query_one(ChatHistory)
# Remove selected class from all citations
for widget in chat_history.query(CitationWidget):
widget.remove_class("selected")
# Add selected class to the newly selected citation
citation_widgets = list(chat_history.query(CitationWidget))
if 0 <= event.citation_index < len(citation_widgets):
citation_widgets[event.citation_index].add_class("selected")
self._selected_citation_idx = event.citation_index