Unify get_typed() into get() with optional type parameter
Fix find_document regression
This commit is contained in:
parent
0dd181356e
commit
da47e3e345
9 changed files with 66 additions and 37 deletions
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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.
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue