Decouple QA toolset from chat agent, add ToolContext state methods
This commit is contained in:
parent
0c6d73409d
commit
8f180f707b
8 changed files with 204 additions and 145 deletions
|
|
@ -1,38 +0,0 @@
|
||||||
from haiku.rag.agents.chat import (
|
|
||||||
ChatDeps,
|
|
||||||
ChatSessionState,
|
|
||||||
create_chat_agent,
|
|
||||||
)
|
|
||||||
from haiku.rag.agents.qa import QuestionAnswerAgent, get_qa_agent
|
|
||||||
from haiku.rag.agents.research import (
|
|
||||||
Citation,
|
|
||||||
IterativePlanResult,
|
|
||||||
ResearchContext,
|
|
||||||
ResearchDependencies,
|
|
||||||
ResearchReport,
|
|
||||||
SearchAnswer,
|
|
||||||
)
|
|
||||||
from haiku.rag.agents.research.graph import build_research_graph
|
|
||||||
from haiku.rag.agents.research.state import ResearchDeps, ResearchState
|
|
||||||
from haiku.rag.tools.qa import QAHistoryEntry
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
# QA
|
|
||||||
"get_qa_agent",
|
|
||||||
"QuestionAnswerAgent",
|
|
||||||
# Research
|
|
||||||
"build_research_graph",
|
|
||||||
"ResearchContext",
|
|
||||||
"ResearchDependencies",
|
|
||||||
"ResearchDeps",
|
|
||||||
"ResearchState",
|
|
||||||
"ResearchReport",
|
|
||||||
"Citation",
|
|
||||||
"SearchAnswer",
|
|
||||||
"IterativePlanResult",
|
|
||||||
# Chat
|
|
||||||
"create_chat_agent",
|
|
||||||
"ChatDeps",
|
|
||||||
"ChatSessionState",
|
|
||||||
"QAHistoryEntry",
|
|
||||||
]
|
|
||||||
|
|
@ -14,7 +14,11 @@ from haiku.rag.agents.chat.prompts import build_chat_prompt
|
||||||
from haiku.rag.agents.chat.state import (
|
from haiku.rag.agents.chat.state import (
|
||||||
AGUI_STATE_KEY,
|
AGUI_STATE_KEY,
|
||||||
ChatSessionState,
|
ChatSessionState,
|
||||||
|
_rebuild_models,
|
||||||
)
|
)
|
||||||
|
from haiku.rag.tools.qa import QAHistoryEntry
|
||||||
|
|
||||||
|
_rebuild_models(QAHistoryEntry)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"AGUI_STATE_KEY",
|
"AGUI_STATE_KEY",
|
||||||
|
|
|
||||||
|
|
@ -7,11 +7,7 @@ from haiku.rag.agents.chat.context import (
|
||||||
trigger_background_summarization as _trigger_summarization,
|
trigger_background_summarization as _trigger_summarization,
|
||||||
)
|
)
|
||||||
from haiku.rag.agents.chat.prompts import build_chat_prompt
|
from haiku.rag.agents.chat.prompts import build_chat_prompt
|
||||||
from haiku.rag.agents.chat.state import (
|
from haiku.rag.agents.chat.state import AGUI_STATE_KEY
|
||||||
AGUI_STATE_KEY,
|
|
||||||
ChatSessionState,
|
|
||||||
build_chat_state_snapshot,
|
|
||||||
)
|
|
||||||
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
|
||||||
|
|
@ -33,6 +29,10 @@ FEATURE_ANALYSIS = "analysis"
|
||||||
DEFAULT_FEATURES = [FEATURE_SEARCH, FEATURE_DOCUMENTS, FEATURE_QA]
|
DEFAULT_FEATURES = [FEATURE_SEARCH, FEATURE_DOCUMENTS, FEATURE_QA]
|
||||||
|
|
||||||
|
|
||||||
|
def _on_qa_complete(qa_session_state: QASessionState, config: AppConfig) -> None:
|
||||||
|
_trigger_summarization(qa_session_state=qa_session_state, config=config)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ChatDeps:
|
class ChatDeps:
|
||||||
"""Dependencies for chat agent.
|
"""Dependencies for chat agent.
|
||||||
|
|
@ -49,12 +49,10 @@ class ChatDeps:
|
||||||
def state(self) -> dict[str, Any]:
|
def state(self) -> dict[str, Any]:
|
||||||
"""Get current state for AG-UI protocol.
|
"""Get current state for AG-UI protocol.
|
||||||
|
|
||||||
Combines SessionState and QASessionState into a single state dict
|
Combines all registered namespace states into a single flat dict,
|
||||||
matching the ChatSessionState schema expected by AG-UI clients.
|
matching the ChatSessionState schema expected by AG-UI clients.
|
||||||
"""
|
"""
|
||||||
session_state = self.tool_context.get(SESSION_NAMESPACE, SessionState)
|
snapshot = self.tool_context.build_state_snapshot()
|
||||||
qa_session_state = self.tool_context.get(QA_SESSION_NAMESPACE, QASessionState)
|
|
||||||
snapshot = build_chat_state_snapshot(session_state, qa_session_state)
|
|
||||||
if self.state_key:
|
if self.state_key:
|
||||||
return {self.state_key: snapshot}
|
return {self.state_key: snapshot}
|
||||||
return snapshot
|
return snapshot
|
||||||
|
|
@ -72,41 +70,26 @@ class ChatDeps:
|
||||||
if isinstance(nested, dict):
|
if isinstance(nested, dict):
|
||||||
state_data = nested
|
state_data = nested
|
||||||
|
|
||||||
session_state = self.tool_context.get(SESSION_NAMESPACE, SessionState)
|
# Preserve server's session_context before restore overwrites it
|
||||||
if session_state is not None:
|
|
||||||
if "document_filter" in state_data:
|
|
||||||
session_state.document_filter = state_data.get("document_filter", [])
|
|
||||||
if "citation_registry" in state_data:
|
|
||||||
session_state.citation_registry = state_data["citation_registry"]
|
|
||||||
if "citations" in state_data:
|
|
||||||
from haiku.rag.agents.research.models import Citation
|
|
||||||
|
|
||||||
session_state.citations = [
|
|
||||||
Citation(**c) if isinstance(c, dict) else c
|
|
||||||
for c in state_data.get("citations", [])
|
|
||||||
]
|
|
||||||
|
|
||||||
qa_session_state = self.tool_context.get(QA_SESSION_NAMESPACE, QASessionState)
|
qa_session_state = self.tool_context.get(QA_SESSION_NAMESPACE, QASessionState)
|
||||||
|
server_session_context = (
|
||||||
|
qa_session_state.session_context if qa_session_state is not None else None
|
||||||
|
)
|
||||||
|
|
||||||
|
self.tool_context.restore_state_snapshot(state_data)
|
||||||
|
|
||||||
|
# Chat-specific overrides after generic restore
|
||||||
if qa_session_state is not None:
|
if qa_session_state is not None:
|
||||||
if "qa_history" in state_data:
|
|
||||||
from haiku.rag.tools.qa import QAHistoryEntry
|
|
||||||
|
|
||||||
qa_session_state.qa_history = [
|
|
||||||
QAHistoryEntry(**qa) if isinstance(qa, dict) else qa
|
|
||||||
for qa in state_data.get("qa_history", [])
|
|
||||||
]
|
|
||||||
|
|
||||||
# Prefer server's session_context (background summarizer may
|
# Prefer server's session_context (background summarizer may
|
||||||
# have updated it since the client's last snapshot).
|
# have updated it since the client's last snapshot).
|
||||||
if qa_session_state.session_context is None:
|
if server_session_context is not None:
|
||||||
session_context = state_data.get("session_context")
|
qa_session_state.session_context = server_session_context
|
||||||
if isinstance(session_context, dict):
|
|
||||||
qa_session_state.session_context = SessionContext(**session_context)
|
|
||||||
|
|
||||||
# Handle initial_context -> session_context for first message
|
# Handle initial_context -> session_context for first message
|
||||||
|
if qa_session_state.session_context is None:
|
||||||
if "initial_context" in state_data:
|
if "initial_context" in state_data:
|
||||||
initial = state_data.get("initial_context")
|
initial = state_data.get("initial_context")
|
||||||
if initial and qa_session_state.session_context is None:
|
if initial:
|
||||||
qa_session_state.session_context = SessionContext(
|
qa_session_state.session_context = SessionContext(
|
||||||
summary=initial
|
summary=initial
|
||||||
)
|
)
|
||||||
|
|
@ -169,7 +152,7 @@ def create_chat_agent(
|
||||||
if FEATURE_DOCUMENTS in features:
|
if FEATURE_DOCUMENTS in features:
|
||||||
toolsets.append(create_document_toolset(config))
|
toolsets.append(create_document_toolset(config))
|
||||||
if FEATURE_QA in features:
|
if FEATURE_QA in features:
|
||||||
toolsets.append(create_qa_toolset(config))
|
toolsets.append(create_qa_toolset(config, on_ask_complete=_on_qa_complete))
|
||||||
if FEATURE_ANALYSIS in features:
|
if FEATURE_ANALYSIS in features:
|
||||||
from haiku.rag.tools.analysis import create_analysis_toolset
|
from haiku.rag.tools.analysis import create_analysis_toolset
|
||||||
|
|
||||||
|
|
@ -234,7 +217,6 @@ __all__ = [
|
||||||
"run_chat_agent",
|
"run_chat_agent",
|
||||||
"trigger_background_summarization",
|
"trigger_background_summarization",
|
||||||
"ChatDeps",
|
"ChatDeps",
|
||||||
"ChatSessionState",
|
|
||||||
"AGUI_STATE_KEY",
|
"AGUI_STATE_KEY",
|
||||||
"FEATURE_SEARCH",
|
"FEATURE_SEARCH",
|
||||||
"FEATURE_DOCUMENTS",
|
"FEATURE_DOCUMENTS",
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
|
@ -6,8 +6,7 @@ from haiku.rag.agents.research.models import Citation
|
||||||
from haiku.rag.tools.session import SessionContext
|
from haiku.rag.tools.session import SessionContext
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from haiku.rag.tools.qa import QAHistoryEntry, QASessionState
|
from haiku.rag.tools.qa import QAHistoryEntry
|
||||||
from haiku.rag.tools.session import SessionState
|
|
||||||
|
|
||||||
AGUI_STATE_KEY = "haiku.rag.chat"
|
AGUI_STATE_KEY = "haiku.rag.chat"
|
||||||
|
|
||||||
|
|
@ -31,39 +30,3 @@ def _rebuild_models(qa_history_entry_cls: type) -> None:
|
||||||
ChatSessionState.model_rebuild(
|
ChatSessionState.model_rebuild(
|
||||||
_types_namespace={"QAHistoryEntry": qa_history_entry_cls}
|
_types_namespace={"QAHistoryEntry": qa_history_entry_cls}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def build_chat_state_snapshot(
|
|
||||||
session_state: "SessionState | None",
|
|
||||||
qa_state: "QASessionState | None",
|
|
||||||
) -> dict[str, Any]:
|
|
||||||
"""Build a combined AG-UI chat state snapshot from current values.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
session_state: SessionState from ToolContext.
|
|
||||||
qa_state: QASessionState from ToolContext.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Snapshot dict.
|
|
||||||
"""
|
|
||||||
snapshot: dict[str, Any] = {}
|
|
||||||
|
|
||||||
if session_state is not None:
|
|
||||||
snapshot.update(
|
|
||||||
{
|
|
||||||
"document_filter": session_state.document_filter.copy(),
|
|
||||||
"citation_registry": session_state.citation_registry.copy(),
|
|
||||||
"citations": [c.model_dump() for c in session_state.citations],
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
if qa_state is not None:
|
|
||||||
snapshot["qa_history"] = [qa.model_dump() for qa in qa_state.qa_history]
|
|
||||||
if qa_state.session_context is not None:
|
|
||||||
snapshot["session_context"] = qa_state.session_context.model_dump(
|
|
||||||
mode="json"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
snapshot["session_context"] = None
|
|
||||||
|
|
||||||
return snapshot
|
|
||||||
|
|
|
||||||
|
|
@ -129,6 +129,41 @@ class ToolContext(BaseModel):
|
||||||
"""
|
"""
|
||||||
return {ns: state.model_dump() for ns, state in self._namespaces.items()}
|
return {ns: state.model_dump() for ns, state in self._namespaces.items()}
|
||||||
|
|
||||||
|
def build_state_snapshot(self) -> dict[str, Any]:
|
||||||
|
"""Build a flat snapshot of all namespace states for AG-UI.
|
||||||
|
|
||||||
|
Merges model_dump(mode="json") from every registered namespace
|
||||||
|
into a single flat dict.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Combined dict of all namespace fields.
|
||||||
|
"""
|
||||||
|
snapshot: dict[str, Any] = {}
|
||||||
|
for state in self._namespaces.values():
|
||||||
|
snapshot.update(state.model_dump(mode="json"))
|
||||||
|
return snapshot
|
||||||
|
|
||||||
|
def restore_state_snapshot(self, data: dict[str, Any]) -> None:
|
||||||
|
"""Restore namespace states from a flat snapshot dict.
|
||||||
|
|
||||||
|
For each registered namespace, finds matching fields in *data*,
|
||||||
|
validates them via the namespace model, and updates the state
|
||||||
|
in place. Fields not present in *data* are left unchanged.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
data: Flat dict as produced by build_state_snapshot().
|
||||||
|
"""
|
||||||
|
for state in self._namespaces.values():
|
||||||
|
model_fields = state.model_fields
|
||||||
|
matching = {k: v for k, v in data.items() if k in model_fields}
|
||||||
|
if matching:
|
||||||
|
# Fill in current values for fields not in data
|
||||||
|
current = state.model_dump()
|
||||||
|
current.update(matching)
|
||||||
|
updated = state.model_validate(current)
|
||||||
|
for field_name in matching:
|
||||||
|
setattr(state, field_name, getattr(updated, field_name))
|
||||||
|
|
||||||
def load_namespace(self, namespace: str, state_type: type[T], data: dict) -> T:
|
def load_namespace(self, namespace: str, state_type: type[T], data: dict) -> T:
|
||||||
"""Deserialize and register state for a namespace.
|
"""Deserialize and register state for a namespace.
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,5 @@
|
||||||
import math
|
import math
|
||||||
|
from collections.abc import Callable
|
||||||
|
|
||||||
from ag_ui.core import EventType, StateSnapshotEvent
|
from ag_ui.core import EventType, StateSnapshotEvent
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
@ -64,15 +65,6 @@ class QAHistoryEntry(BaseModel):
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _resolve_chat_state_forward_refs() -> None:
|
|
||||||
from haiku.rag.agents.chat.state import _rebuild_models
|
|
||||||
|
|
||||||
_rebuild_models(QAHistoryEntry)
|
|
||||||
|
|
||||||
|
|
||||||
_resolve_chat_state_forward_refs()
|
|
||||||
|
|
||||||
|
|
||||||
class QASessionState(BaseModel):
|
class QASessionState(BaseModel):
|
||||||
"""Extended session state for QA with embedding cache."""
|
"""Extended session state for QA with embedding cache."""
|
||||||
|
|
||||||
|
|
@ -94,6 +86,7 @@ async def run_qa_core(
|
||||||
base_filter: str | None = None,
|
base_filter: str | None = None,
|
||||||
session_context: str | None = None,
|
session_context: str | None = None,
|
||||||
prior_answers: list[SearchAnswer] | None = None,
|
prior_answers: list[SearchAnswer] | None = None,
|
||||||
|
on_qa_complete: Callable[[QASessionState, AppConfig], None] | None = None,
|
||||||
) -> QAResult:
|
) -> QAResult:
|
||||||
"""Run the QA flow and return a QAResult.
|
"""Run the QA flow and return a QAResult.
|
||||||
|
|
||||||
|
|
@ -205,12 +198,8 @@ async def run_qa_core(
|
||||||
# Enforce FIFO limit
|
# Enforce FIFO limit
|
||||||
if len(qa_session_state.qa_history) > MAX_QA_HISTORY:
|
if len(qa_session_state.qa_history) > MAX_QA_HISTORY:
|
||||||
qa_session_state.qa_history = qa_session_state.qa_history[-MAX_QA_HISTORY:]
|
qa_session_state.qa_history = qa_session_state.qa_history[-MAX_QA_HISTORY:]
|
||||||
from haiku.rag.agents.chat.context import trigger_background_summarization
|
if on_qa_complete is not None:
|
||||||
|
on_qa_complete(qa_session_state, config)
|
||||||
trigger_background_summarization(
|
|
||||||
qa_session_state=qa_session_state,
|
|
||||||
config=config,
|
|
||||||
)
|
|
||||||
|
|
||||||
return qa_result
|
return qa_result
|
||||||
|
|
||||||
|
|
@ -219,6 +208,7 @@ def create_qa_toolset(
|
||||||
config: AppConfig,
|
config: AppConfig,
|
||||||
base_filter: str | None = None,
|
base_filter: str | None = None,
|
||||||
tool_name: str = "ask",
|
tool_name: str = "ask",
|
||||||
|
on_ask_complete: Callable[[QASessionState, AppConfig], None] | None = None,
|
||||||
) -> FunctionToolset:
|
) -> FunctionToolset:
|
||||||
"""Create a toolset with Q&A capabilities using research graph.
|
"""Create a toolset with Q&A capabilities using research graph.
|
||||||
|
|
||||||
|
|
@ -226,6 +216,9 @@ def create_qa_toolset(
|
||||||
config: Application configuration.
|
config: Application configuration.
|
||||||
base_filter: Optional base SQL WHERE clause applied to searches.
|
base_filter: Optional base SQL WHERE clause applied to searches.
|
||||||
tool_name: Name for the ask tool. Defaults to "ask".
|
tool_name: Name for the ask tool. Defaults to "ask".
|
||||||
|
on_ask_complete: Optional callback invoked after each QA cycle with
|
||||||
|
the updated QASessionState and config. Use this to trigger
|
||||||
|
background summarization or other post-processing.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
FunctionToolset with an ask tool.
|
FunctionToolset with an ask tool.
|
||||||
|
|
@ -250,13 +243,9 @@ def create_qa_toolset(
|
||||||
client = ctx.deps.client
|
client = ctx.deps.client
|
||||||
tool_context = ctx.deps.tool_context
|
tool_context = ctx.deps.tool_context
|
||||||
|
|
||||||
session_state: SessionState | None = None
|
|
||||||
qa_session_state: QASessionState | None = None
|
|
||||||
state_key: str | None = None
|
state_key: str | None = None
|
||||||
|
|
||||||
if tool_context is not None:
|
if tool_context is not None:
|
||||||
session_state = tool_context.get(SESSION_NAMESPACE, SessionState)
|
|
||||||
qa_session_state = tool_context.get(QA_SESSION_NAMESPACE, QASessionState)
|
|
||||||
state_key = tool_context.state_key
|
state_key = tool_context.state_key
|
||||||
|
|
||||||
qa_result = await run_qa_core(
|
qa_result = await run_qa_core(
|
||||||
|
|
@ -266,15 +255,11 @@ def create_qa_toolset(
|
||||||
document_name=document_name,
|
document_name=document_name,
|
||||||
context=tool_context,
|
context=tool_context,
|
||||||
base_filter=base_filter,
|
base_filter=base_filter,
|
||||||
|
on_qa_complete=on_ask_complete,
|
||||||
)
|
)
|
||||||
|
|
||||||
if session_state is not None:
|
if tool_context is not None and tool_context.namespaces:
|
||||||
from haiku.rag.agents.chat.state import build_chat_state_snapshot
|
snapshot = tool_context.build_state_snapshot()
|
||||||
|
|
||||||
snapshot = build_chat_state_snapshot(
|
|
||||||
session_state,
|
|
||||||
qa_session_state,
|
|
||||||
)
|
|
||||||
if state_key:
|
if state_key:
|
||||||
snapshot = {state_key: snapshot}
|
snapshot = {state_key: snapshot}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -561,7 +561,7 @@ async def test_chat_agent_ask_triggers_background_summarization(
|
||||||
)
|
)
|
||||||
|
|
||||||
# Patch internal trigger to avoid concurrent HTTP calls during VCR
|
# Patch internal trigger to avoid concurrent HTTP calls during VCR
|
||||||
with patch("haiku.rag.agents.chat.context.trigger_background_summarization"):
|
with patch("haiku.rag.agents.chat.agent._trigger_summarization"):
|
||||||
result = await agent.run(
|
result = await agent.run(
|
||||||
"What is the highest count class in the DocLayNet dataset?",
|
"What is the highest count class in the DocLayNet dataset?",
|
||||||
deps=deps,
|
deps=deps,
|
||||||
|
|
@ -650,7 +650,7 @@ async def test_chat_agent_multi_turn_with_context(allow_model_requests, temp_db_
|
||||||
# Patch the internal summarization trigger in the ask tool to avoid
|
# Patch the internal summarization trigger in the ask tool to avoid
|
||||||
# concurrent HTTP calls that break VCR cassette replay ordering.
|
# concurrent HTTP calls that break VCR cassette replay ordering.
|
||||||
with patch(
|
with patch(
|
||||||
"haiku.rag.agents.chat.context.trigger_background_summarization",
|
"haiku.rag.agents.chat.agent._trigger_summarization",
|
||||||
):
|
):
|
||||||
# First question about class labels
|
# First question about class labels
|
||||||
result1 = await agent.run(
|
result1 = await agent.run(
|
||||||
|
|
@ -675,7 +675,7 @@ async def test_chat_agent_multi_turn_with_context(allow_model_requests, temp_db_
|
||||||
|
|
||||||
# Second related question - uses prior answers and updated session context
|
# Second related question - uses prior answers and updated session context
|
||||||
with patch(
|
with patch(
|
||||||
"haiku.rag.agents.chat.context.trigger_background_summarization",
|
"haiku.rag.agents.chat.agent._trigger_summarization",
|
||||||
):
|
):
|
||||||
result2 = await agent.run(
|
result2 = await agent.run(
|
||||||
"How were the annotations created and how many annotators were involved?",
|
"How were the annotations created and how many annotators were involved?",
|
||||||
|
|
|
||||||
|
|
@ -312,3 +312,131 @@ def test_tool_context_cache_clear():
|
||||||
ctx2, is_new2 = cache.get_or_create("thread-2")
|
ctx2, is_new2 = cache.get_or_create("thread-2")
|
||||||
assert is_new1 is True
|
assert is_new1 is True
|
||||||
assert is_new2 is True
|
assert is_new2 is True
|
||||||
|
|
||||||
|
|
||||||
|
# =============================================================================
|
||||||
|
# build_state_snapshot / restore_state_snapshot Tests
|
||||||
|
# =============================================================================
|
||||||
|
|
||||||
|
|
||||||
|
class NestedModel(BaseModel):
|
||||||
|
name: str = ""
|
||||||
|
count: int = 0
|
||||||
|
|
||||||
|
|
||||||
|
class StateWithNested(BaseModel):
|
||||||
|
nested: NestedModel | None = None
|
||||||
|
tags: list[str] = []
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_state_snapshot_empty():
|
||||||
|
"""build_state_snapshot on empty context returns empty dict."""
|
||||||
|
ctx = ToolContext()
|
||||||
|
assert ctx.build_state_snapshot() == {}
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_state_snapshot_single_namespace():
|
||||||
|
"""build_state_snapshot with one namespace returns its fields."""
|
||||||
|
ctx = ToolContext()
|
||||||
|
ctx.register("ns1", TestState(value=42))
|
||||||
|
snapshot = ctx.build_state_snapshot()
|
||||||
|
assert snapshot == {"value": 42}
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_state_snapshot_multiple_namespaces():
|
||||||
|
"""build_state_snapshot merges fields from all namespaces."""
|
||||||
|
ctx = ToolContext()
|
||||||
|
ctx.register("ns1", TestState(value=42))
|
||||||
|
ctx.register("ns2", TestStateWithList(items=["a", "b"]))
|
||||||
|
snapshot = ctx.build_state_snapshot()
|
||||||
|
assert snapshot == {"value": 42, "items": ["a", "b"]}
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_state_snapshot_nested_model():
|
||||||
|
"""build_state_snapshot serializes nested models with mode='json'."""
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
class TimestampState(BaseModel):
|
||||||
|
ts: datetime | None = None
|
||||||
|
|
||||||
|
ctx = ToolContext()
|
||||||
|
ctx.register("ns", TimestampState(ts=datetime(2025, 1, 27, 12, 0, 0)))
|
||||||
|
snapshot = ctx.build_state_snapshot()
|
||||||
|
assert isinstance(snapshot["ts"], str)
|
||||||
|
assert snapshot["ts"] == "2025-01-27T12:00:00"
|
||||||
|
|
||||||
|
|
||||||
|
def test_restore_state_snapshot_empty_context():
|
||||||
|
"""restore_state_snapshot on empty context is a no-op."""
|
||||||
|
ctx = ToolContext()
|
||||||
|
ctx.restore_state_snapshot({"value": 42})
|
||||||
|
assert ctx.namespaces == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_restore_state_snapshot_single_namespace():
|
||||||
|
"""restore_state_snapshot updates matching fields in registered namespaces."""
|
||||||
|
ctx = ToolContext()
|
||||||
|
ctx.register("ns1", TestState(value=0))
|
||||||
|
ctx.restore_state_snapshot({"value": 99})
|
||||||
|
state = ctx.get("ns1", TestState)
|
||||||
|
assert state is not None
|
||||||
|
assert state.value == 99
|
||||||
|
|
||||||
|
|
||||||
|
def test_restore_state_snapshot_partial_update():
|
||||||
|
"""restore_state_snapshot only touches fields present in data."""
|
||||||
|
ctx = ToolContext()
|
||||||
|
ctx.register("ns1", TestState(value=42))
|
||||||
|
ctx.register("ns2", TestStateWithList(items=["original"]))
|
||||||
|
# Only update ns2's items, not ns1's value
|
||||||
|
ctx.restore_state_snapshot({"items": ["updated"]})
|
||||||
|
ns1 = ctx.get("ns1", TestState)
|
||||||
|
ns2 = ctx.get("ns2", TestStateWithList)
|
||||||
|
assert ns1 is not None
|
||||||
|
assert ns2 is not None
|
||||||
|
assert ns1.value == 42
|
||||||
|
assert ns2.items == ["updated"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_restore_state_snapshot_nested_model():
|
||||||
|
"""restore_state_snapshot deserializes nested models from dicts."""
|
||||||
|
ctx = ToolContext()
|
||||||
|
ctx.register("ns", StateWithNested())
|
||||||
|
ctx.restore_state_snapshot({"nested": {"name": "foo", "count": 5}, "tags": ["x"]})
|
||||||
|
state = ctx.get("ns", StateWithNested)
|
||||||
|
assert state is not None
|
||||||
|
assert state.nested is not None
|
||||||
|
assert state.nested.name == "foo"
|
||||||
|
assert state.nested.count == 5
|
||||||
|
assert state.tags == ["x"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_state_snapshot_roundtrip():
|
||||||
|
"""build then restore produces equivalent state."""
|
||||||
|
ctx = ToolContext()
|
||||||
|
ctx.register("ns1", TestState(value=42))
|
||||||
|
ctx.register("ns2", TestStateWithList(items=["a", "b"]))
|
||||||
|
|
||||||
|
snapshot = ctx.build_state_snapshot()
|
||||||
|
|
||||||
|
ctx2 = ToolContext()
|
||||||
|
ctx2.register("ns1", TestState())
|
||||||
|
ctx2.register("ns2", TestStateWithList())
|
||||||
|
ctx2.restore_state_snapshot(snapshot)
|
||||||
|
|
||||||
|
ns1 = ctx2.get("ns1", TestState)
|
||||||
|
ns2 = ctx2.get("ns2", TestStateWithList)
|
||||||
|
assert ns1 is not None
|
||||||
|
assert ns2 is not None
|
||||||
|
assert ns1.value == 42
|
||||||
|
assert ns2.items == ["a", "b"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_restore_state_snapshot_ignores_unknown_fields():
|
||||||
|
"""restore_state_snapshot ignores fields not in any registered namespace."""
|
||||||
|
ctx = ToolContext()
|
||||||
|
ctx.register("ns1", TestState(value=0))
|
||||||
|
ctx.restore_state_snapshot({"value": 10, "unknown_field": "ignored"})
|
||||||
|
ns1 = ctx.get("ns1", TestState)
|
||||||
|
assert ns1 is not None
|
||||||
|
assert ns1.value == 10
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue