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.agents.chat.state import AGUI_STATE_KEY
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.config.models import AppConfig
|
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.document import create_document_toolset
|
||||||
from haiku.rag.tools.qa import (
|
from haiku.rag.tools.qa import (
|
||||||
QA_SESSION_NAMESPACE,
|
QA_SESSION_NAMESPACE,
|
||||||
|
|
@ -18,7 +18,7 @@ from haiku.rag.tools.qa import (
|
||||||
create_qa_toolset,
|
create_qa_toolset,
|
||||||
)
|
)
|
||||||
from haiku.rag.tools.search import create_search_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
|
from haiku.rag.utils import get_model
|
||||||
|
|
||||||
FEATURE_SEARCH = "search"
|
FEATURE_SEARCH = "search"
|
||||||
|
|
@ -110,14 +110,7 @@ def prepare_chat_context(
|
||||||
if features is None:
|
if features is None:
|
||||||
features = DEFAULT_FEATURES
|
features = DEFAULT_FEATURES
|
||||||
|
|
||||||
if context.get(SESSION_NAMESPACE, SessionState) is None:
|
prepare_context(context, features=features, state_key=AGUI_STATE_KEY)
|
||||||
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())
|
|
||||||
|
|
||||||
|
|
||||||
def create_chat_agent(
|
def create_chat_agent(
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,11 @@
|
||||||
from haiku.rag.tools.analysis import create_analysis_toolset
|
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 (
|
from haiku.rag.tools.document import (
|
||||||
DocumentInfo,
|
DocumentInfo,
|
||||||
DocumentListResponse,
|
DocumentListResponse,
|
||||||
|
|
@ -30,9 +36,11 @@ from haiku.rag.tools.session import (
|
||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
"AgentDeps",
|
||||||
"RAGDeps",
|
"RAGDeps",
|
||||||
"ToolContext",
|
"ToolContext",
|
||||||
"ToolContextCache",
|
"ToolContextCache",
|
||||||
|
"prepare_context",
|
||||||
"QAResult",
|
"QAResult",
|
||||||
"AnalysisResult",
|
"AnalysisResult",
|
||||||
"build_document_filter",
|
"build_document_filter",
|
||||||
|
|
|
||||||
|
|
@ -180,6 +180,36 @@ class ToolContext(BaseModel):
|
||||||
return state
|
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:
|
class ToolContextCache:
|
||||||
"""In-memory cache for ToolContext instances, keyed by external session/thread ID."""
|
"""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 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):
|
class TestState(BaseModel):
|
||||||
|
|
@ -440,3 +442,47 @@ def test_restore_state_snapshot_ignores_unknown_fields():
|
||||||
ns1 = ctx.get("ns1", TestState)
|
ns1 = ctx.get("ns1", TestState)
|
||||||
assert ns1 is not None
|
assert ns1 is not None
|
||||||
assert ns1.value == 10
|
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