Add AgentDeps and prepare_context for custom agent DX

This commit is contained in:
Yiorgis Gozadinos 2026-02-12 17:34:52 +02:00
parent 8f180f707b
commit a5b8e3be7e
No known key found for this signature in database
6 changed files with 229 additions and 12 deletions

View file

@ -10,7 +10,7 @@ from haiku.rag.agents.chat.prompts import build_chat_prompt
from haiku.rag.agents.chat.state import AGUI_STATE_KEY
from haiku.rag.client import HaikuRAG
from haiku.rag.config.models import AppConfig
from haiku.rag.tools.context import ToolContext
from haiku.rag.tools.context import ToolContext, prepare_context
from haiku.rag.tools.document import create_document_toolset
from haiku.rag.tools.qa import (
QA_SESSION_NAMESPACE,
@ -18,7 +18,7 @@ from haiku.rag.tools.qa import (
create_qa_toolset,
)
from haiku.rag.tools.search import create_search_toolset
from haiku.rag.tools.session import SESSION_NAMESPACE, SessionContext, SessionState
from haiku.rag.tools.session import SessionContext
from haiku.rag.utils import get_model
FEATURE_SEARCH = "search"
@ -110,14 +110,7 @@ def prepare_chat_context(
if features is None:
features = DEFAULT_FEATURES
if context.get(SESSION_NAMESPACE, SessionState) is None:
context.register(SESSION_NAMESPACE, SessionState())
if context.state_key is None:
context.state_key = AGUI_STATE_KEY
if FEATURE_QA in features:
if context.get(QA_SESSION_NAMESPACE, QASessionState) is None:
context.register(QA_SESSION_NAMESPACE, QASessionState())
prepare_context(context, features=features, state_key=AGUI_STATE_KEY)
def create_chat_agent(

View file

@ -1,5 +1,11 @@
from haiku.rag.tools.analysis import create_analysis_toolset
from haiku.rag.tools.context import RAGDeps, ToolContext, ToolContextCache
from haiku.rag.tools.context import (
RAGDeps,
ToolContext,
ToolContextCache,
prepare_context,
)
from haiku.rag.tools.deps import AgentDeps
from haiku.rag.tools.document import (
DocumentInfo,
DocumentListResponse,
@ -30,9 +36,11 @@ from haiku.rag.tools.session import (
)
__all__ = [
"AgentDeps",
"RAGDeps",
"ToolContext",
"ToolContextCache",
"prepare_context",
"QAResult",
"AnalysisResult",
"build_document_filter",

View file

@ -180,6 +180,36 @@ class ToolContext(BaseModel):
return state
def prepare_context(
context: ToolContext,
features: list[str] | None = None,
state_key: str | None = None,
) -> None:
"""Register required namespaces in a ToolContext based on feature flags.
Idempotent safe to call multiple times on the same context.
Args:
context: ToolContext to prepare.
features: List of enabled features. Defaults to ["search", "documents"].
state_key: Optional AG-UI state key to set on the context.
"""
from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState
from haiku.rag.tools.session import SESSION_NAMESPACE, SessionState
if features is None:
features = ["search", "documents"]
if any(f in features for f in ("search", "qa", "analysis")):
context.get_or_create(SESSION_NAMESPACE, SessionState)
if "qa" in features:
context.get_or_create(QA_SESSION_NAMESPACE, QASessionState)
if state_key is not None:
context.state_key = state_key
class ToolContextCache:
"""In-memory cache for ToolContext instances, keyed by external session/thread ID."""

View file

@ -0,0 +1,38 @@
from dataclasses import dataclass
from typing import Any
from haiku.rag.client import HaikuRAG
from haiku.rag.tools.context import ToolContext
@dataclass
class AgentDeps:
"""Generic dependencies for agents using haiku.rag toolsets.
Implements RAGDeps protocol and AG-UI state protocol.
"""
client: HaikuRAG
tool_context: ToolContext
state_key: str | None = None
@property
def state(self) -> dict[str, Any]:
"""Get current state for AG-UI protocol."""
snapshot = self.tool_context.build_state_snapshot()
if self.state_key:
return {self.state_key: snapshot}
return snapshot
@state.setter
def state(self, value: dict[str, Any] | None) -> None:
"""Set state from AG-UI protocol."""
if value is None:
return
data: dict[str, Any] = value
if self.state_key and self.state_key in value:
nested = value[self.state_key]
if isinstance(nested, dict):
data = nested
self.tool_context.restore_state_snapshot(data)

View file

@ -1,6 +1,8 @@
from pydantic import BaseModel
from haiku.rag.tools.context import ToolContext
from haiku.rag.tools.context import ToolContext, prepare_context
from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState
from haiku.rag.tools.session import SESSION_NAMESPACE, SessionState
class TestState(BaseModel):
@ -440,3 +442,47 @@ def test_restore_state_snapshot_ignores_unknown_fields():
ns1 = ctx.get("ns1", TestState)
assert ns1 is not None
assert ns1.value == 10
# --- prepare_context tests ---
def test_prepare_context_default_features():
"""Default features register SessionState only."""
ctx = ToolContext()
prepare_context(ctx)
assert ctx.get(SESSION_NAMESPACE, SessionState) is not None
assert ctx.get(QA_SESSION_NAMESPACE, QASessionState) is None
def test_prepare_context_with_qa():
"""QA feature registers both SessionState and QASessionState."""
ctx = ToolContext()
prepare_context(ctx, features=["search", "qa"])
assert ctx.get(SESSION_NAMESPACE, SessionState) is not None
assert ctx.get(QA_SESSION_NAMESPACE, QASessionState) is not None
def test_prepare_context_sets_state_key():
"""state_key is set on context when provided."""
ctx = ToolContext()
prepare_context(ctx, state_key="my_app")
assert ctx.state_key == "my_app"
def test_prepare_context_no_state_key_by_default():
"""state_key is not set when not provided."""
ctx = ToolContext()
prepare_context(ctx)
assert ctx.state_key is None
def test_prepare_context_idempotent():
"""Calling prepare_context twice doesn't create duplicate state."""
ctx = ToolContext()
prepare_context(ctx, features=["search", "qa"])
session1 = ctx.get(SESSION_NAMESPACE, SessionState)
qa1 = ctx.get(QA_SESSION_NAMESPACE, QASessionState)
prepare_context(ctx, features=["search", "qa"])
assert ctx.get(SESSION_NAMESPACE, SessionState) is session1
assert ctx.get(QA_SESSION_NAMESPACE, QASessionState) is qa1

102
tests/tools/test_deps.py Normal file
View file

@ -0,0 +1,102 @@
from unittest.mock import MagicMock
import pytest
from haiku.rag.tools.context import ToolContext
from haiku.rag.tools.deps import AgentDeps
from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState
from haiku.rag.tools.session import SESSION_NAMESPACE, SessionState
@pytest.fixture
def mock_client():
return MagicMock()
def test_agent_deps_state_getter_empty(mock_client):
"""state returns empty dict when no namespaces are registered."""
ctx = ToolContext()
deps = AgentDeps(client=mock_client, tool_context=ctx)
assert deps.state == {}
def test_agent_deps_state_getter_with_session(mock_client):
"""state returns flat snapshot of registered namespaces."""
ctx = ToolContext()
ctx.register(SESSION_NAMESPACE, SessionState())
deps = AgentDeps(client=mock_client, tool_context=ctx)
snapshot = deps.state
assert "citations" in snapshot
assert "citation_registry" in snapshot
def test_agent_deps_state_getter_with_state_key(mock_client):
"""state wraps snapshot under state_key when set."""
ctx = ToolContext()
ctx.register(SESSION_NAMESPACE, SessionState())
deps = AgentDeps(client=mock_client, tool_context=ctx, state_key="my_app")
snapshot = deps.state
assert "my_app" in snapshot
assert "citations" in snapshot["my_app"]
def test_agent_deps_state_setter_restores(mock_client):
"""state setter restores namespace fields from flat dict."""
ctx = ToolContext()
ctx.register(SESSION_NAMESPACE, SessionState())
deps = AgentDeps(client=mock_client, tool_context=ctx)
deps.state = {"document_filter": ["doc1", "doc2"]}
session = ctx.get(SESSION_NAMESPACE, SessionState)
assert session is not None
assert session.document_filter == ["doc1", "doc2"]
def test_agent_deps_state_setter_with_state_key(mock_client):
"""state setter extracts data from namespaced key."""
ctx = ToolContext()
ctx.register(SESSION_NAMESPACE, SessionState())
deps = AgentDeps(client=mock_client, tool_context=ctx, state_key="my_app")
deps.state = {"my_app": {"document_filter": ["doc1"]}}
session = ctx.get(SESSION_NAMESPACE, SessionState)
assert session is not None
assert session.document_filter == ["doc1"]
def test_agent_deps_state_setter_ignores_none(mock_client):
"""state setter is a no-op when value is None."""
ctx = ToolContext()
ctx.register(SESSION_NAMESPACE, SessionState())
deps = AgentDeps(client=mock_client, tool_context=ctx)
deps.state = None
session = ctx.get(SESSION_NAMESPACE, SessionState)
assert session is not None
assert session.document_filter == []
def test_agent_deps_state_roundtrip(mock_client):
"""Build snapshot then restore produces equivalent state."""
ctx = ToolContext()
ctx.register(SESSION_NAMESPACE, SessionState(document_filter=["doc1"]))
ctx.register(QA_SESSION_NAMESPACE, QASessionState())
deps = AgentDeps(client=mock_client, tool_context=ctx, state_key="app")
snapshot = deps.state
ctx2 = ToolContext()
ctx2.register(SESSION_NAMESPACE, SessionState())
ctx2.register(QA_SESSION_NAMESPACE, QASessionState())
deps2 = AgentDeps(client=mock_client, tool_context=ctx2, state_key="app")
deps2.state = snapshot
session = ctx2.get(SESSION_NAMESPACE, SessionState)
assert session is not None
assert session.document_filter == ["doc1"]
def test_agent_deps_satisfies_rag_deps_protocol(mock_client):
"""AgentDeps satisfies the RAGDeps protocol."""
from haiku.rag.tools.context import RAGDeps
ctx = ToolContext()
deps = AgentDeps(client=mock_client, tool_context=ctx)
assert isinstance(deps, RAGDeps)