Unify get_typed() into get() with optional type parameter

Fix find_document regression
This commit is contained in:
Yiorgis Gozadinos 2026-02-10 11:10:09 +02:00
parent 0dd181356e
commit da47e3e345
No known key found for this signature in database
9 changed files with 66 additions and 37 deletions

View file

@ -52,16 +52,14 @@ class ChatDeps:
snapshot: dict[str, Any] = {"session_id": self.session_id} snapshot: dict[str, Any] = {"session_id": self.session_id}
# Add SessionState fields # Add SessionState fields
session_state = self.tool_context.get_typed(SESSION_NAMESPACE, SessionState) session_state = self.tool_context.get(SESSION_NAMESPACE, SessionState)
if session_state is not None: if session_state is not None:
snapshot["document_filter"] = session_state.document_filter snapshot["document_filter"] = session_state.document_filter
snapshot["citation_registry"] = session_state.citation_registry snapshot["citation_registry"] = session_state.citation_registry
snapshot["citations"] = [c.model_dump() for c in session_state.citations] snapshot["citations"] = [c.model_dump() for c in session_state.citations]
# Add QASessionState fields # Add QASessionState fields
qa_session_state = self.tool_context.get_typed( qa_session_state = self.tool_context.get(QA_SESSION_NAMESPACE, QASessionState)
QA_SESSION_NAMESPACE, QASessionState
)
if qa_session_state is not None: if qa_session_state is not None:
snapshot["qa_history"] = [ snapshot["qa_history"] = [
qa.model_dump() for qa in qa_session_state.qa_history qa.model_dump() for qa in qa_session_state.qa_history
@ -92,7 +90,7 @@ class ChatDeps:
state_data = nested state_data = nested
# Update SessionState from incoming state # Update SessionState from incoming state
session_state = self.tool_context.get_typed(SESSION_NAMESPACE, SessionState) session_state = self.tool_context.get(SESSION_NAMESPACE, SessionState)
if session_state is not None: if session_state is not None:
if "document_filter" in state_data: if "document_filter" in state_data:
session_state.document_filter = state_data.get("document_filter", []) session_state.document_filter = state_data.get("document_filter", [])
@ -122,9 +120,7 @@ class ChatDeps:
session_state.incoming_session_id = incoming_session_id session_state.incoming_session_id = incoming_session_id
# Update QASessionState from incoming state # Update QASessionState from incoming state
qa_session_state = self.tool_context.get_typed( qa_session_state = self.tool_context.get(QA_SESSION_NAMESPACE, QASessionState)
QA_SESSION_NAMESPACE, QASessionState
)
if qa_session_state is not None: if qa_session_state is not None:
if "qa_history" in state_data: if "qa_history" in state_data:
from haiku.rag.tools.qa import QAHistoryEntry from haiku.rag.tools.qa import QAHistoryEntry
@ -187,12 +183,12 @@ def create_chat_agent(
result = await agent.run("Search for X", deps=deps) result = await agent.run("Search for X", deps=deps)
""" """
# Ensure session states are registered with proper AG-UI state key # Ensure session states are registered with proper AG-UI state key
existing = context.get_typed(SESSION_NAMESPACE, SessionState) existing = context.get(SESSION_NAMESPACE, SessionState)
if existing is None: if existing is None:
context.register(SESSION_NAMESPACE, SessionState(state_key=AGUI_STATE_KEY)) context.register(SESSION_NAMESPACE, SessionState(state_key=AGUI_STATE_KEY))
elif existing.state_key is None: elif existing.state_key is None:
existing.state_key = AGUI_STATE_KEY existing.state_key = AGUI_STATE_KEY
if context.get_typed(QA_SESSION_NAMESPACE, QASessionState) is None: if context.get(QA_SESSION_NAMESPACE, QASessionState) is None:
context.register(QA_SESSION_NAMESPACE, QASessionState()) context.register(QA_SESSION_NAMESPACE, QASessionState())
# Create toolsets - these capture client, config, and context in closures # Create toolsets - these capture client, config, and context in closures
@ -230,7 +226,7 @@ def trigger_background_summarization(deps: ChatDeps) -> None:
Args: Args:
deps: Chat dependencies with tool_context containing QASessionState. deps: Chat dependencies with tool_context containing QASessionState.
""" """
qa_session_state = deps.tool_context.get_typed(QA_SESSION_NAMESPACE, QASessionState) qa_session_state = deps.tool_context.get(QA_SESSION_NAMESPACE, QASessionState)
if qa_session_state is None or not qa_session_state.qa_history: if qa_session_state is None or not qa_session_state.qa_history:
return return
if not deps.session_id: if not deps.session_id:

View file

@ -1,7 +1,7 @@
from datetime import datetime from datetime import datetime
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from pydantic import BaseModel, Field from pydantic import BaseModel
from haiku.rag.agents.research.models import Citation from haiku.rag.agents.research.models import Citation

View file

@ -163,7 +163,7 @@ class ChatApp(App):
self.agent = create_chat_agent(self.config, self.client, self.tool_context) self.agent = create_chat_agent(self.config, self.client, self.tool_context)
# Initialize session state in tool context # Initialize session state in tool context
session_state = self.tool_context.get_typed(SESSION_NAMESPACE, SessionState) session_state = self.tool_context.get(SESSION_NAMESPACE, SessionState)
if session_state is not None: if session_state is not None:
session_state.document_filter = self._document_filter session_state.document_filter = self._document_filter
@ -320,7 +320,7 @@ class ChatApp(App):
} }
# Sync session state to tool context before running # Sync session state to tool context before running
session_state = self.tool_context.get_typed(SESSION_NAMESPACE, SessionState) session_state = self.tool_context.get(SESSION_NAMESPACE, SessionState)
if session_state is not None: if session_state is not None:
session_state.session_id = self.session_state.session_id session_state.session_id = self.session_state.session_id
session_state.document_filter = self.session_state.document_filter session_state.document_filter = self.session_state.document_filter
@ -330,7 +330,7 @@ class ChatApp(App):
# Sync initial_context to QA session state # Sync initial_context to QA session state
from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState
qa_session_state = self.tool_context.get_typed( qa_session_state = self.tool_context.get(
QA_SESSION_NAMESPACE, QASessionState QA_SESSION_NAMESPACE, QASessionState
) )
if qa_session_state is not None: if qa_session_state is not None:
@ -379,7 +379,7 @@ class ChatApp(App):
trigger_background_summarization(deps) trigger_background_summarization(deps)
# Sync session context from QASessionState to ChatSessionState for modal # Sync session context from QASessionState to ChatSessionState for modal
qa_session_state = self.tool_context.get_typed( qa_session_state = self.tool_context.get(
QA_SESSION_NAMESPACE, QASessionState QA_SESSION_NAMESPACE, QASessionState
) )
if qa_session_state is not None and qa_session_state.session_context: if qa_session_state is not None and qa_session_state.session_context:
@ -471,9 +471,7 @@ class ChatApp(App):
from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState
# Sync session context from QASessionState before showing modal # Sync session context from QASessionState before showing modal
qa_session_state = self.tool_context.get_typed( qa_session_state = self.tool_context.get(QA_SESSION_NAMESPACE, QASessionState)
QA_SESSION_NAMESPACE, QASessionState
)
if qa_session_state is not None and qa_session_state.session_context: if qa_session_state is not None and qa_session_state.session_context:
from datetime import datetime from datetime import datetime

View file

@ -1,4 +1,4 @@
from typing import Any, TypeVar from typing import Any, TypeVar, overload
from pydantic import BaseModel, PrivateAttr from pydantic import BaseModel, PrivateAttr
@ -57,19 +57,24 @@ class ToolContext(BaseModel):
""" """
self._namespaces[namespace] = state self._namespaces[namespace] = state
def get(self, namespace: str) -> BaseModel | None: @overload
"""Get state for a namespace, or None if not registered.""" def get(self, namespace: str) -> BaseModel | None: ...
return self._namespaces.get(namespace)
def get_typed(self, namespace: str, expected_type: type[T]) -> T | None: @overload
"""Get state for a namespace with type checking. def get(self, namespace: str, state_type: type[T]) -> T | None: ...
Returns the state cast to expected_type if it matches, None otherwise. def get(
self, namespace: str, state_type: type[T] | None = None
) -> BaseModel | T | None:
"""Get state for a namespace, or None if not registered.
When state_type is provided, returns the state only if it matches
the expected type, otherwise returns None.
""" """
state = self._namespaces.get(namespace) state = self._namespaces.get(namespace)
if isinstance(state, expected_type): if state_type is not None:
return state return state if isinstance(state, state_type) else None
return None return state
def get_or_create(self, namespace: str, state_type: type[T]) -> T: def get_or_create(self, namespace: str, state_type: type[T]) -> T:
"""Get state for a namespace, creating it if not registered. """Get state for a namespace, creating it if not registered.

View file

@ -65,16 +65,16 @@ async def find_document(client: HaikuRAG, query: str):
limit=1, limit=1,
filter=f"LOWER(uri) LIKE LOWER('%{escaped_query}%') OR LOWER(uri) LIKE LOWER('%{no_spaces}%')", filter=f"LOWER(uri) LIKE LOWER('%{escaped_query}%') OR LOWER(uri) LIKE LOWER('%{no_spaces}%')",
) )
if docs: if docs and docs[0].id:
return docs[0] return await client.get_document_by_id(docs[0].id)
# Try partial title match (with and without spaces) # Try partial title match (with and without spaces)
docs = await client.list_documents( docs = await client.list_documents(
limit=1, limit=1,
filter=f"LOWER(title) LIKE LOWER('%{escaped_query}%') OR LOWER(title) LIKE LOWER('%{no_spaces}%')", filter=f"LOWER(title) LIKE LOWER('%{escaped_query}%') OR LOWER(title) LIKE LOWER('%{no_spaces}%')",
) )
if docs: if docs and docs[0].id:
return docs[0] return await client.get_document_by_id(docs[0].id)
return None return None

View file

@ -52,7 +52,7 @@ def get_session_filter(
from haiku.rag.tools.session import SESSION_NAMESPACE, SessionState from haiku.rag.tools.session import SESSION_NAMESPACE, SessionState
session_state = context.get_typed(SESSION_NAMESPACE, SessionState) session_state = context.get(SESSION_NAMESPACE, SessionState)
if session_state is None or not session_state.document_filter: if session_state is None or not session_state.document_filter:
return base_filter return base_filter

View file

@ -154,8 +154,8 @@ def create_qa_toolset(
old_state_snapshot: dict | None = None old_state_snapshot: dict | None = None
if context is not None: if context is not None:
session_state = context.get_typed(SESSION_NAMESPACE, SessionState) session_state = context.get(SESSION_NAMESPACE, SessionState)
qa_session_state = context.get_typed(QA_SESSION_NAMESPACE, QASessionState) qa_session_state = context.get(QA_SESSION_NAMESPACE, QASessionState)
# Capture combined state snapshot before changes # Capture combined state snapshot before changes
# Use incoming values (what client sent) so delta shows server-side updates # Use incoming values (what client sent) so delta shows server-side updates

View file

@ -71,7 +71,7 @@ def create_search_toolset(
session_state: SessionState | None = None session_state: SessionState | None = None
old_session_state: SessionState | None = None old_session_state: SessionState | None = None
if context is not None: if context is not None:
session_state = context.get_typed(SESSION_NAMESPACE, SessionState) session_state = context.get(SESSION_NAMESPACE, SessionState)
if session_state is not None: if session_state is not None:
old_session_state = session_state.model_copy(deep=True) old_session_state = session_state.model_copy(deep=True)

View file

@ -171,3 +171,33 @@ def test_serialization_roundtrip():
qa_state = restored.get("qa") qa_state = restored.get("qa")
assert isinstance(qa_state, TestState) assert isinstance(qa_state, TestState)
assert qa_state.value == 99 assert qa_state.value == 99
def test_get_with_type_match():
"""Test get with state_type returns typed state when type matches."""
ctx = ToolContext()
state = TestState(value=42)
ctx.register("ns", state)
result = ctx.get("ns", TestState)
assert result is state
assert result.value == 42
def test_get_with_type_mismatch():
"""Test get with state_type returns None when type doesn't match."""
ctx = ToolContext()
ctx.register("ns", TestState(value=42))
result = ctx.get("ns", TestStateWithList)
assert result is None
def test_get_without_type():
"""Test get without state_type returns BaseModel (unchanged behavior)."""
ctx = ToolContext()
state = TestState(value=42)
ctx.register("ns", state)
result = ctx.get("ns")
assert result is state