Replace console/stream with AG-UI event protocol

This commit is contained in:
Yiorgis Gozadinos 2025-11-10 15:08:33 +02:00
parent da77e0f2f9
commit 663a7fac9a
No known key found for this signature in database
11 changed files with 502 additions and 451 deletions

View file

@ -1,5 +1,6 @@
"""Generic AG-UI protocol support for haiku.rag graphs."""
from haiku.rag.agui.cli_renderer import AGUIConsoleRenderer
from haiku.rag.agui.emitter import AGUIEmitter
from haiku.rag.agui.events import (
AGUIEvent,
@ -18,8 +19,10 @@ from haiku.rag.agui.events import (
emit_text_message_start,
)
from haiku.rag.agui.state import compute_state_delta
from haiku.rag.agui.stream import stream_graph
__all__ = [
"AGUIConsoleRenderer",
"AGUIEmitter",
"AGUIEvent",
"compute_state_delta",
@ -36,4 +39,5 @@ __all__ = [
"emit_text_message_content",
"emit_text_message_end",
"emit_text_message_start",
"stream_graph",
]

View file

@ -0,0 +1,152 @@
"""Generic CLI renderer for AG-UI events with Rich console output."""
from collections.abc import AsyncIterator
from typing import Any
from pydantic import BaseModel
from rich.console import Console
from haiku.rag.agui.events import AGUIEvent
class AGUIConsoleRenderer:
"""Renders AG-UI events to Rich console with formatted output.
Generic renderer that processes AG-UI protocol events and renders them
with Rich formatting. Works with any graph that emits AG-UI events.
"""
def __init__(self, console: Console | None = None):
"""Initialize the renderer.
Args:
console: Optional Rich console instance (creates new one if not provided)
"""
self.console = console or Console()
self._state: BaseModel | None = None
async def render(self, events: AsyncIterator[AGUIEvent]) -> Any | None:
"""Process events and render to console, return final result.
Args:
events: Async iterator of AG-UI events
Returns:
The final result from RunFinished event, or None
"""
result = None
async for event in events:
event_type = event.get("type")
if event_type == "RUN_STARTED":
self._render_run_started(event)
elif event_type == "RUN_FINISHED":
result = event.get("result")
self._render_run_finished()
elif event_type == "RUN_ERROR":
self._render_error(event)
elif event_type == "STEP_STARTED":
self._render_step_started(event)
elif event_type == "STEP_FINISHED":
self._render_step_finished(event)
elif event_type == "TEXT_MESSAGE_CHUNK":
self._render_text_message(event)
elif event_type == "TEXT_MESSAGE_START":
pass # Start of streaming message, no output needed
elif event_type == "TEXT_MESSAGE_CONTENT":
self._render_text_content(event)
elif event_type == "TEXT_MESSAGE_END":
pass # End of streaming message, no output needed
elif event_type == "STATE_SNAPSHOT":
self._state = event.get("snapshot")
elif event_type == "STATE_DELTA":
self._apply_state_delta(event)
elif event_type == "ACTIVITY_SNAPSHOT":
self._render_activity(event)
elif event_type == "ACTIVITY_DELTA":
pass # Activity deltas don't need separate rendering
return result
def _render_run_started(self, event: AGUIEvent) -> None:
"""Render run start event.
Args:
event: RunStarted event
"""
# Currently silent - could render run metadata later
def _render_run_finished(self) -> None:
"""Render run completion."""
# Currently silent - the result is rendered separately
def _render_error(self, event: AGUIEvent) -> None:
"""Render error event.
Args:
event: RunError event
"""
message = event.get("message", "Unknown error")
self.console.print(f"[bold red]❌ Error:[/bold red] {message}")
def _render_step_started(self, event: AGUIEvent) -> None:
"""Render step start event.
Args:
event: StepStarted event
"""
step_name = event.get("stepName", "")
if step_name:
# Format step name for display
display_name = step_name.replace("_", " ").title()
self.console.print(f"\n[bold cyan]{display_name}[/bold cyan]")
def _render_step_finished(self, event: AGUIEvent) -> None:
"""Render step finish event.
Args:
event: StepFinished event
"""
# Step completion is implicit from the next step or activity
def _render_text_message(self, event: AGUIEvent) -> None:
"""Render complete text message.
Args:
event: TextMessageChunk event
"""
delta = event.get("delta", "")
# The delta contains the text content to display
self.console.print(delta)
def _render_text_content(self, event: AGUIEvent) -> None:
"""Render streaming text content delta.
Args:
event: TextMessageContent event
"""
delta = event.get("delta", "")
# Print delta without newline for streaming effect
self.console.print(delta, end="")
def _render_activity(self, event: AGUIEvent) -> None:
"""Render activity update.
Args:
event: ActivitySnapshot event
"""
content = event.get("content", "")
# Render activity content with emphasis
if content:
self.console.print(f"[dim]{content}[/dim]")
def _apply_state_delta(self, event: AGUIEvent) -> None:
"""Apply state delta to current state.
Args:
event: StateDelta event
"""
# Currently not applying deltas - could implement state patching later
# For now, the text messages contain all the information we need to display

View file

@ -13,7 +13,6 @@ from haiku.rag.agui.events import (
emit_run_error,
emit_run_finished,
emit_run_started,
emit_state_delta,
emit_state_snapshot,
emit_step_finished,
emit_step_started,
@ -60,18 +59,19 @@ class AGUIEmitter[StateT: BaseModel, ResultT]:
"""Get the run ID for this emitter."""
return self._run_id
def start_run(self, input_data: str, initial_state: StateT) -> None:
def start_run(self, initial_state: StateT) -> None:
"""Emit RunStarted and initial StateSnapshot.
Args:
input_data: The input that started the run
initial_state: The initial state of the graph
"""
# If thread_id wasn't provided, generate from input hash
# If thread_id wasn't provided, generate from state hash
if not self._thread_id or self._thread_id == str(uuid4()):
self._thread_id = self._generate_thread_id(input_data)
state_json = initial_state.model_dump_json()
self._thread_id = self._generate_thread_id(state_json)
self._emit(emit_run_started(self._thread_id, self._run_id, input_data))
# RunStarted (state snapshot follows immediately with full state)
self._emit(emit_run_started(self._thread_id, self._run_id))
self._emit(emit_state_snapshot(initial_state))
self._last_state = initial_state
@ -100,21 +100,13 @@ class AGUIEmitter[StateT: BaseModel, ResultT]:
self._emit(emit_text_message(message, role))
def update_state(self, new_state: StateT) -> None:
"""Emit StateDelta for state change, or StateSnapshot if no previous state.
"""Emit StateSnapshot for state change.
Args:
new_state: The updated state
"""
if self._last_state:
# Emit delta if we have a previous state
delta_event = emit_state_delta(self._last_state, new_state)
# Only emit if there are actual changes
if delta_event.get("operations"):
self._emit(delta_event)
else:
# Emit snapshot if this is the first state update
# Always emit full snapshot (not delta) for complete state visibility
self._emit(emit_state_snapshot(new_state))
self._last_state = new_state
def update_activity(
@ -125,9 +117,11 @@ class AGUIEmitter[StateT: BaseModel, ResultT]:
Args:
activity_type: Type of activity (e.g., "planning", "searching")
content: Description of the activity
message_id: Optional message ID to associate activity with
message_id: Optional message ID to associate activity with (auto-generated if None)
"""
self._emit(emit_activity(activity_type, content, message_id))
if message_id is None:
message_id = str(uuid4())
self._emit(emit_activity(message_id, activity_type, content))
def finish_run(self, result: ResultT) -> None:
"""Emit RunFinished event.

View file

@ -11,23 +11,27 @@ from haiku.rag.agui.state import compute_state_delta
AGUIEvent = dict[str, Any]
def emit_run_started(thread_id: str, run_id: str, input_data: str) -> dict[str, Any]:
def emit_run_started(
thread_id: str, run_id: str, input_data: str | None = None
) -> dict[str, Any]:
"""Create a RunStarted event.
Args:
thread_id: Unique identifier for the conversation thread
run_id: Unique identifier for this run
input_data: The input that started the run
input_data: Optional input that started the run
Returns:
RunStarted event dict
"""
return {
event: dict[str, Any] = {
"type": "RUN_STARTED",
"threadId": thread_id,
"runId": run_id,
"input": input_data,
}
if input_data:
event["input"] = input_data
return event
def emit_run_finished(thread_id: str, run_id: str, result: Any) -> dict[str, Any]:
@ -119,7 +123,7 @@ def emit_text_message(content: str, role: str = "assistant") -> dict[str, Any]:
"type": "TEXT_MESSAGE_CHUNK",
"messageId": message_id,
"role": role,
"content": content,
"delta": content,
}
@ -197,46 +201,47 @@ def emit_state_delta(old_state: BaseModel, new_state: BaseModel) -> dict[str, An
Returns:
StateDelta event dict
"""
operations = compute_state_delta(old_state, new_state)
delta = compute_state_delta(old_state, new_state)
return {
"type": "STATE_DELTA",
"operations": operations,
"delta": delta,
}
def emit_activity(
message_id: str,
activity_type: str,
content: str,
message_id: str | None = None,
) -> dict[str, Any]:
"""Create an ActivitySnapshot event.
Args:
message_id: Message ID to associate activity with (required)
activity_type: Type of activity (e.g., "planning", "searching")
content: Description of the activity
message_id: Optional message ID to associate activity with
Returns:
ActivitySnapshot event dict
"""
event: dict[str, Any] = {
return {
"type": "ACTIVITY_SNAPSHOT",
"messageId": message_id,
"activityType": activity_type,
"content": content,
}
if message_id:
event["messageId"] = message_id
return event
def emit_activity_delta(
message_id: str, operations: list[dict[str, Any]]
message_id: str,
activity_type: str,
patch: list[dict[str, Any]],
) -> dict[str, Any]:
"""Create an ActivityDelta event with JSON Patch operations.
Args:
message_id: Message ID of the activity being updated
operations: JSON Patch operations to apply
activity_type: Type of activity being updated
patch: JSON Patch operations to apply
Returns:
ActivityDelta event dict
@ -244,5 +249,6 @@ def emit_activity_delta(
return {
"type": "ACTIVITY_DELTA",
"messageId": message_id,
"operations": operations,
"activityType": activity_type,
"patch": patch,
}

View file

@ -0,0 +1,80 @@
"""Generic graph streaming with AG-UI events."""
import asyncio
from collections.abc import AsyncIterator
from contextlib import suppress
from typing import Any, Protocol
from pydantic import BaseModel
from haiku.rag.agui.emitter import AGUIEmitter
from haiku.rag.agui.events import AGUIEvent
class GraphDeps(Protocol):
"""Protocol for graph dependencies that support AG-UI emission."""
agui_emitter: AGUIEmitter[Any, Any] | None
async def stream_graph(
graph: Any,
state: BaseModel,
deps: GraphDeps,
) -> AsyncIterator[AGUIEvent]:
"""Run a graph and yield AG-UI events as they occur.
This is a generic streaming function that works with any pydantic-graph
that follows the AG-UI pattern:
- State must be a Pydantic BaseModel
- Deps must have an optional agui_emitter attribute
- Graph must be a pydantic-graph Graph instance
Args:
graph: The pydantic-graph Graph to execute
state: Initial state (Pydantic BaseModel)
deps: Graph dependencies with agui_emitter support
Yields:
AG-UI event dictionaries
Raises:
TypeError: If deps doesn't support agui_emitter
RuntimeError: If graph doesn't produce a result
"""
if not hasattr(deps, "agui_emitter"):
raise TypeError("deps must have an 'agui_emitter' attribute")
# Create AG-UI emitter
emitter: AGUIEmitter[Any, Any] = AGUIEmitter()
deps.agui_emitter = emitter
async def _execute() -> None:
try:
# Start the run with initial state
emitter.start_run(initial_state=state)
# Execute the graph
result = await graph.run(state=state, deps=deps)
if result is None:
raise RuntimeError("Graph did not produce a result")
# Finish the run with the result
emitter.finish_run(result)
except Exception as exc:
# Emit error event
emitter.error(exc)
finally:
await emitter.close()
runner = asyncio.create_task(_execute())
try:
async for event in emitter:
yield event
finally:
if not runner.done():
runner.cancel()
with suppress(asyncio.CancelledError):
await runner

View file

@ -8,6 +8,7 @@ from rich.console import Console
from rich.markdown import Markdown
from rich.progress import Progress
from haiku.rag.agui import AGUIConsoleRenderer, stream_graph
from haiku.rag.client import HaikuRAG
from haiku.rag.config import AppConfig, Config
from haiku.rag.mcp import create_mcp_server
@ -15,7 +16,6 @@ from haiku.rag.monitor import FileWatcher
from haiku.rag.research.dependencies import ResearchContext
from haiku.rag.research.graph import build_research_graph
from haiku.rag.research.state import ResearchDeps, ResearchState
from haiku.rag.research.stream import stream_research_graph
from haiku.rag.store.models.chunk import Chunk
from haiku.rag.store.models.document import Document
@ -243,20 +243,14 @@ class HaikuRAGApp:
except Exception as e:
self.console.print(f"[red]Error: {e}[/red]")
async def research(
self,
question: str,
verbose: bool = False,
):
async def research(self, question: str):
"""Run research via the pydantic-graph pipeline.
Args:
question: The research question
verbose: Show verbose output
"""
async with HaikuRAG(db_path=self.db_path, config=self.config) as client:
try:
if verbose:
self.console.print("[bold cyan]Starting research[/bold cyan]")
self.console.print(f"[bold blue]Question:[/bold blue] {question}")
self.console.print()
@ -264,25 +258,21 @@ class HaikuRAGApp:
graph = build_research_graph(config=self.config)
context = ResearchContext(original_question=question)
state = ResearchState.from_config(context=context, config=self.config)
deps = ResearchDeps(
client=client, console=self.console if verbose else None
)
deps = ResearchDeps(client=client)
report = None
async for event in stream_research_graph(graph, state, deps):
if event.type == "report":
report = event.report
break
if event.type == "error":
self.console.print(
f"[red]Error during research: {event.message}[/red]"
)
return
# Use AG-UI renderer to process events
renderer = AGUIConsoleRenderer(self.console)
report_dict = await renderer.render(stream_graph(graph, state, deps))
if report is None:
if report_dict is None:
self.console.print("[red]Research did not produce a report.[/red]")
return
# Convert dict to ResearchReport model
from haiku.rag.research.models import ResearchReport
report = ResearchReport.model_validate(report_dict)
# Display the report
self.console.print("[bold green]Research Report[/bold green]")
self.console.rule()

View file

@ -299,14 +299,9 @@ def research(
"--db",
help="Path to the LanceDB database file",
),
verbose: bool = typer.Option(
False,
"--verbose",
help="Show verbose progress output",
),
):
app = create_app(db)
asyncio.run(app.research(question=question, verbose=verbose))
asyncio.run(app.research(question=question))
@cli.command("settings", help="Display current configuration settings")

View file

@ -1,7 +1,6 @@
from collections.abc import Iterable
from pydantic import BaseModel, Field
from rich.console import Console
from haiku.rag.client import HaikuRAG
from haiku.rag.graph_common.models import SearchAnswer
@ -10,7 +9,6 @@ from haiku.rag.research.models import (
InsightAnalysis,
InsightRecord,
)
from haiku.rag.research.stream import ResearchStream
class ResearchContext(BaseModel):
@ -184,10 +182,6 @@ class ResearchDependencies(BaseModel):
client: HaikuRAG = Field(description="RAG client for document operations")
context: ResearchContext = Field(description="Shared research context")
console: Console | None = None
stream: ResearchStream | None = Field(
default=None, description="Optional research event stream"
)
def _merge_unique(existing: list[str], incoming: Iterable[str]) -> list[str]:

View file

@ -53,8 +53,11 @@ def build_research_graph(
state = ctx.state
deps = ctx.deps
log(deps, state, "\n[bold cyan]📋 Creating research plan...[/bold cyan]")
if deps.agui_emitter:
deps.agui_emitter.start_step("plan")
deps.agui_emitter.update_activity("planning", "Creating research plan")
try:
plan_agent = Agent(
model=get_model(provider, model),
output_type=ResearchPlan,
@ -63,6 +66,7 @@ def build_research_graph(
+ "\n\nUse the gather_context tool once on the main question before planning."
),
retries=3,
output_retries=3,
deps_type=ResearchDependencies,
)
@ -82,21 +86,18 @@ def build_research_graph(
agent_deps = ResearchDependencies(
client=deps.client,
context=state.context,
console=deps.console,
stream=deps.stream,
)
plan_result = await plan_agent.run(prompt, deps=agent_deps)
state.context.sub_questions = list(plan_result.output.sub_questions)
log(deps, state, "\n[bold green]✅ Plan Created:[/bold green]")
log(
deps,
state,
f" [bold]Main Question:[/bold] {state.context.original_question}",
)
log(deps, state, " [bold]Sub-questions:[/bold]")
# Log the plan results
log(deps, state, f"Main Question: {state.context.original_question}")
log(deps, state, "Sub-questions:")
for i, sq in enumerate(state.context.sub_questions, 1):
log(deps, state, f" {i}. {sq}")
finally:
if deps.agui_emitter:
deps.agui_emitter.finish_step()
@g.step
async def search_one(
@ -106,6 +107,10 @@ def build_research_graph(
deps = ctx.deps
sub_q = ctx.inputs
if deps.agui_emitter:
deps.agui_emitter.start_step("search_one")
try:
# Create semaphore if not already provided
if deps.semaphore is None:
import asyncio
@ -115,17 +120,17 @@ def build_research_graph(
# Use semaphore to control concurrency
async with deps.semaphore:
return await _do_search(state, deps, sub_q)
finally:
if deps.agui_emitter:
deps.agui_emitter.finish_step()
async def _do_search(
state: ResearchState,
deps: ResearchDeps,
sub_q: str,
) -> SearchAnswer:
log(
deps,
state,
f"\n[bold cyan]🔍 Searching & Answering:[/bold cyan] {sub_q}",
)
if deps.agui_emitter:
deps.agui_emitter.update_activity("searching", f"Searching: {sub_q}")
agent = Agent(
model=get_model(provider, model),
@ -160,21 +165,16 @@ def build_research_graph(
agent_deps = ResearchDependencies(
client=deps.client,
context=state.context,
console=deps.console,
stream=deps.stream,
)
try:
result = await agent.run(sub_q, deps=agent_deps)
answer = result.output
if answer:
state.context.add_qa_response(answer)
preview = answer.answer[:150] + (
"" if len(answer.answer) > 150 else ""
)
log(deps, state, f" [green]✓[/green] {preview}")
log(deps, state, f"Answer: {answer.answer}")
return answer
except Exception as e:
log(deps, state, f"[red]Search failed:[/red] {e}")
log(deps, state, f"Search failed: {e}")
failure_answer = SearchAnswer(
query=sub_q,
answer=f"Search failed after retries: {str(e)}",
@ -204,17 +204,19 @@ def build_research_graph(
state = ctx.state
deps = ctx.deps
log(
deps,
state,
"\n[bold cyan]🧭 Synthesizing new insights and gap status...[/bold cyan]",
if deps.agui_emitter:
deps.agui_emitter.start_step("analyze_insights")
deps.agui_emitter.update_activity(
"analyzing", "Synthesizing insights and gaps"
)
try:
agent = Agent(
model=get_model(provider, model),
output_type=InsightAnalysis,
instructions=INSIGHT_AGENT_PROMPT,
retries=3,
output_retries=3,
deps_type=ResearchDependencies,
)
@ -227,8 +229,6 @@ def build_research_graph(
agent_deps = ResearchDependencies(
client=deps.client,
context=state.context,
console=deps.console,
stream=deps.stream,
)
result = await agent.run(prompt, deps=agent_deps)
analysis: InsightAnalysis = result.output
@ -239,49 +239,46 @@ def build_research_graph(
if analysis.commentary:
log(deps, state, f"Summary: {analysis.commentary}")
if analysis.highlights:
log(deps, state, " [bold]Updated insights:[/bold]")
log(deps, state, "Updated insights:")
for insight in analysis.highlights:
label = insight.status.value
log(
deps,
state,
f" • ({label}) {insight.summary}",
)
log(deps, state, f" • ({label}) {insight.summary}")
if analysis.gap_assessments:
log(deps, state, " [bold yellow]Gap updates:[/bold yellow]")
log(deps, state, "Gap updates:")
for gap in analysis.gap_assessments:
status = "resolved" if gap.resolved else "open"
severity = gap.severity.value
log(
deps,
state,
f" • ({severity}/{status}) {gap.description}",
)
log(deps, state, f" • ({severity}/{status}) {gap.description}")
if analysis.resolved_gaps:
log(deps, state, " [green]Resolved gaps:[/green]")
log(deps, state, "Resolved gaps:")
for resolved in analysis.resolved_gaps:
log(deps, state, f"{resolved}")
if analysis.new_questions:
log(deps, state, " [cyan]Proposed follow-ups:[/cyan]")
log(deps, state, "Proposed follow-ups:")
for question in analysis.new_questions:
log(deps, state, f"{question}")
finally:
if deps.agui_emitter:
deps.agui_emitter.finish_step()
@g.step
async def decide(ctx: StepContext[ResearchState, ResearchDeps, None]) -> bool:
state = ctx.state
deps = ctx.deps
log(
deps,
state,
"\n[bold cyan]📊 Evaluating research sufficiency...[/bold cyan]",
if deps.agui_emitter:
deps.agui_emitter.start_step("decide")
deps.agui_emitter.update_activity(
"evaluating", "Evaluating research sufficiency"
)
try:
agent = Agent(
model=get_model(provider, model),
output_type=EvaluationResult,
instructions=DECISION_AGENT_PROMPT,
retries=3,
output_retries=3,
deps_type=ResearchDependencies,
)
@ -306,8 +303,6 @@ def build_research_graph(
agent_deps = ResearchDependencies(
client=deps.client,
context=state.context,
console=deps.console,
stream=deps.stream,
)
decision_result = await agent.run(prompt, deps=agent_deps)
output = decision_result.output
@ -320,21 +315,17 @@ def build_research_graph(
state.context.sub_questions.append(new_q)
if output.key_insights:
log(deps, state, " [bold]Key insights:[/bold]")
log(deps, state, "Key insights:")
for insight in output.key_insights:
log(deps, state, f"{insight}")
if output.gaps:
log(deps, state, " [bold yellow]Remaining gaps:[/bold yellow]")
log(deps, state, "Remaining gaps:")
for gap in output.gaps:
log(deps, state, f"{gap}")
log(
deps,
state,
f" Confidence: [yellow]{output.confidence_score:.1%}[/yellow]",
)
status = "[green]Yes[/green]" if output.is_sufficient else "[red]No[/red]"
log(deps, state, f"Confidence: {output.confidence_score:.1%}")
status = "Yes" if output.is_sufficient else "No"
log(deps, state, f"Sufficient: {status}")
should_continue = (
@ -342,10 +333,10 @@ def build_research_graph(
or output.confidence_score < state.confidence_threshold
) and state.iterations < state.max_iterations
if not should_continue:
log(deps, state, "\n[bold green]✅ Stopping research.[/bold green]")
return should_continue
finally:
if deps.agui_emitter:
deps.agui_emitter.finish_step()
@g.step
async def synthesize(
@ -354,17 +345,19 @@ def build_research_graph(
state = ctx.state
deps = ctx.deps
log(
deps,
state,
"\n[bold cyan]📝 Generating final research report...[/bold cyan]",
if deps.agui_emitter:
deps.agui_emitter.start_step("synthesize")
deps.agui_emitter.update_activity(
"synthesizing", "Generating final research report"
)
try:
agent = Agent(
model=get_model(provider, model),
output_type=ResearchReport,
instructions=SYNTHESIS_AGENT_PROMPT,
retries=3,
output_retries=3,
deps_type=ResearchDependencies,
)
@ -377,13 +370,12 @@ def build_research_graph(
agent_deps = ResearchDependencies(
client=deps.client,
context=state.context,
console=deps.console,
stream=deps.stream,
)
result = await agent.run(prompt, deps=agent_deps)
log(deps, state, "[bold green]✅ Research complete![/bold green]")
return result.output
finally:
if deps.agui_emitter:
deps.agui_emitter.finish_step()
# Build the graph structure
collect_answers = g.join(

View file

@ -3,29 +3,35 @@ from dataclasses import dataclass
from typing import TYPE_CHECKING
from pydantic import BaseModel, Field
from rich.console import Console
from haiku.rag.client import HaikuRAG
from haiku.rag.research.dependencies import ResearchContext
from haiku.rag.research.models import EvaluationResult, InsightAnalysis
from haiku.rag.research.stream import ResearchStream
from haiku.rag.research.models import EvaluationResult, InsightAnalysis, ResearchReport
if TYPE_CHECKING:
from haiku.rag.agui.emitter import AGUIEmitter
from haiku.rag.config.models import AppConfig
@dataclass
class ResearchDeps:
"""Dependencies for research graph execution."""
client: HaikuRAG
console: Console | None = None
stream: ResearchStream | None = None
agui_emitter: "AGUIEmitter[ResearchState, ResearchReport] | None" = None
semaphore: asyncio.Semaphore | None = None
def emit_log(self, message: str, state: "ResearchState | None" = None) -> None:
if self.console:
self.console.print(message)
if self.stream:
self.stream.log(message, state)
"""Emit a log message through AG-UI events.
Args:
message: The message to log
state: Optional state to include in state update
"""
if self.agui_emitter:
self.agui_emitter.log(message)
if state:
self.agui_emitter.update_state(state)
class ResearchState(BaseModel):

View file

@ -1,162 +0,0 @@
import asyncio
from collections.abc import AsyncIterator
from dataclasses import dataclass
from typing import TYPE_CHECKING, Literal
from haiku.rag.research.models import ResearchReport
if TYPE_CHECKING: # pragma: no cover
from haiku.rag.research.state import ResearchState
@dataclass(slots=True)
class ResearchStateSnapshot:
question: str
sub_questions: list[str]
iterations: int
max_iterations: int
confidence_threshold: float
pending_sub_questions: int
answered_questions: int
insights: list[str]
gaps: list[str]
last_confidence: float | None
last_sufficient: bool | None
@classmethod
def from_state(cls, state: "ResearchState") -> "ResearchStateSnapshot":
context = state.context
last_confidence: float | None = None
last_sufficient: bool | None = None
if state.last_eval:
last_confidence = state.last_eval.confidence_score
last_sufficient = state.last_eval.is_sufficient
return cls(
question=context.original_question,
sub_questions=list(context.sub_questions),
iterations=state.iterations,
max_iterations=state.max_iterations,
confidence_threshold=state.confidence_threshold,
pending_sub_questions=len(context.sub_questions),
answered_questions=len(context.qa_responses),
insights=[
f"{insight.status.value}:{insight.summary}"
for insight in context.insights
],
gaps=[
f"{gap.severity.value}/{'resolved' if gap.resolved else 'open'}:{gap.description}"
for gap in context.gaps
],
last_confidence=last_confidence,
last_sufficient=last_sufficient,
)
@dataclass(slots=True)
class ResearchStreamEvent:
type: Literal["log", "report", "error"]
message: str | None = None
state: ResearchStateSnapshot | None = None
report: ResearchReport | None = None
error: str | None = None
class ResearchStream:
"""Queue-backed stream for research graph events."""
def __init__(self) -> None:
self._queue: asyncio.Queue[ResearchStreamEvent | None] = asyncio.Queue()
self._closed = False
def _snapshot(self, state: "ResearchState | None") -> ResearchStateSnapshot | None:
if state is None:
return None
return ResearchStateSnapshot.from_state(state)
def log(self, message: str, state: "ResearchState | None" = None) -> None:
if self._closed:
return
event = ResearchStreamEvent(
type="log", message=message, state=self._snapshot(state)
)
self._queue.put_nowait(event)
def report(self, report: ResearchReport, state: "ResearchState") -> None:
if self._closed:
return
event = ResearchStreamEvent(
type="report",
report=report,
state=self._snapshot(state),
)
self._queue.put_nowait(event)
def error(self, error: Exception, state: "ResearchState | None" = None) -> None:
if self._closed:
return
event = ResearchStreamEvent(
type="error",
message=str(error),
error=str(error),
state=self._snapshot(state),
)
self._queue.put_nowait(event)
async def close(self) -> None:
if self._closed:
return
self._closed = True
await self._queue.put(None)
def __aiter__(self) -> AsyncIterator[ResearchStreamEvent]:
return self._iter_events()
async def _iter_events(self) -> AsyncIterator[ResearchStreamEvent]:
while True:
event = await self._queue.get()
if event is None:
break
yield event
async def stream_research_graph(
graph,
state: "ResearchState",
deps,
) -> AsyncIterator[ResearchStreamEvent]:
"""Run the research graph and yield streaming events as they occur."""
from contextlib import suppress
from haiku.rag.research.state import ResearchDeps
if not isinstance(deps, ResearchDeps):
raise TypeError("deps must be an instance of ResearchDeps")
stream = ResearchStream()
deps.stream = stream
async def _execute() -> None:
try:
report = await graph.run(state=state, deps=deps)
if report is None:
raise RuntimeError("Graph did not produce a report")
stream.report(report, state)
except Exception as exc:
stream.error(exc, state)
finally:
await stream.close()
runner = asyncio.create_task(_execute())
try:
async for event in stream:
yield event
finally:
if not runner.done():
runner.cancel()
with suppress(asyncio.CancelledError):
await runner