Add AgentDeps and prepare_context for custom agent DX
This commit is contained in:
parent
8f180f707b
commit
a5b8e3be7e
6 changed files with 229 additions and 12 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
38
haiku_rag_slim/haiku/rag/tools/deps.py
Normal file
38
haiku_rag_slim/haiku/rag/tools/deps.py
Normal 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)
|
||||
|
|
@ -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
102
tests/tools/test_deps.py
Normal 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)
|
||||
Loading…
Reference in a new issue