Clean up server vs client session setting priorities
This commit is contained in:
parent
0434cd068b
commit
58c2bd49ea
7 changed files with 84 additions and 49 deletions
|
|
@ -79,13 +79,12 @@ async def stream_chat(request: Request) -> Response:
|
||||||
run_input = AGUIAdapter.build_run_input(body)
|
run_input = AGUIAdapter.build_run_input(body)
|
||||||
|
|
||||||
thread_id = getattr(run_input, "thread_id", None) or "default"
|
thread_id = getattr(run_input, "thread_id", None) or "default"
|
||||||
context, is_new = context_cache.get_or_create(thread_id)
|
context, _is_new = context_cache.get_or_create(thread_id)
|
||||||
agent = create_chat_agent(Config, get_client(), context)
|
agent = create_chat_agent(Config, get_client(), context)
|
||||||
|
|
||||||
deps = ChatDeps(
|
deps = ChatDeps(
|
||||||
config=Config,
|
config=Config,
|
||||||
tool_context=context,
|
tool_context=context,
|
||||||
is_new=is_new,
|
|
||||||
state_key=AGUI_STATE_KEY,
|
state_key=AGUI_STATE_KEY,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -127,13 +127,27 @@ The system prompt is automatically composed to match the selected features. See
|
||||||
|
|
||||||
### Session State
|
### Session State
|
||||||
|
|
||||||
The `ChatSessionState` maintains:
|
Session state is managed through `ToolContext` — a namespace-based state container shared across all toolsets. The chat agent uses two namespaces:
|
||||||
|
|
||||||
|
**`SessionState`** (session management):
|
||||||
|
|
||||||
- `session_id` — Unique identifier for the session
|
|
||||||
- `qa_history` — List of previous Q/A pairs
|
|
||||||
- `session_context` — Automatically maintained session context summary
|
|
||||||
- `document_filter` — List of document titles/URIs to restrict searches
|
- `document_filter` — List of document titles/URIs to restrict searches
|
||||||
- `citation_registry` — Stable mapping of chunk IDs to citation indices
|
- `citation_registry` — Stable mapping of chunk IDs to citation indices
|
||||||
|
- `citations` — Citations from the current query
|
||||||
|
|
||||||
|
**`QASessionState`** (QA history and context):
|
||||||
|
|
||||||
|
- `qa_history` — List of previous Q/A pairs with embeddings
|
||||||
|
- `session_context` — Automatically maintained session context summary
|
||||||
|
|
||||||
|
For multi-session applications (e.g., web backends), use `ToolContextCache` to cache `ToolContext` instances by external session/thread ID:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from haiku.rag.tools import ToolContext, ToolContextCache
|
||||||
|
|
||||||
|
cache = ToolContextCache() # TTL-based, defaults to 1 hour
|
||||||
|
context, _is_new = cache.get_or_create(thread_id)
|
||||||
|
```
|
||||||
|
|
||||||
**Citation Registry**: Citation indices persist across tool calls within a session. The same `chunk_id` always returns the same citation index (first-occurrence-wins). This ensures consistent citation numbering in multi-turn conversations — `[1]` always refers to the same source.
|
**Citation Registry**: Citation indices persist across tool calls within a session. The same `chunk_id` always returns the same citation index (first-occurrence-wins). This ensures consistent citation numbering in multi-turn conversations — `[1]` always refers to the same source.
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -227,9 +227,12 @@ When using the chat agent with [AG-UI](https://docs.ag-ui.com) streaming, `ChatD
|
||||||
|
|
||||||
```python
|
```python
|
||||||
from haiku.rag.agents.chat import AGUI_STATE_KEY, ChatDeps, create_chat_agent
|
from haiku.rag.agents.chat import AGUI_STATE_KEY, ChatDeps, create_chat_agent
|
||||||
from haiku.rag.tools import ToolContext
|
from haiku.rag.tools import ToolContext, ToolContextCache
|
||||||
|
|
||||||
|
# For multi-session apps, cache ToolContext per thread
|
||||||
|
cache = ToolContextCache()
|
||||||
|
context, _is_new = cache.get_or_create(thread_id)
|
||||||
|
|
||||||
context = ToolContext()
|
|
||||||
agent = create_chat_agent(config, client, context)
|
agent = create_chat_agent(config, client, context)
|
||||||
deps = ChatDeps(
|
deps = ChatDeps(
|
||||||
config=config,
|
config=config,
|
||||||
|
|
@ -243,7 +246,6 @@ The emitted state structure:
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"haiku.rag.chat": {
|
"haiku.rag.chat": {
|
||||||
"session_id": "uuid",
|
|
||||||
"citations": [],
|
"citations": [],
|
||||||
"qa_history": [],
|
"qa_history": [],
|
||||||
"session_context": null,
|
"session_context": null,
|
||||||
|
|
@ -253,7 +255,7 @@ The emitted state structure:
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
State flows bidirectionally — the frontend sends its current state on each request, and the agent emits deltas (JSON Patch) reflecting server-side updates (new citations, QA history entries, session context). See the [Web Application](apps.md#web-application) for a complete implementation.
|
State flows bidirectionally — the frontend sends its current state on each request, and the agent emits deltas (JSON Patch) reflecting server-side updates (new citations, QA history entries, session context). The server always prefers its own `session_context` over the client's value, since background summarization may have updated it between requests. See the [Web Application](apps.md#web-application) for a complete implementation.
|
||||||
|
|
||||||
## Filter Helpers
|
## Filter Helpers
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -43,7 +43,6 @@ class ChatDeps:
|
||||||
|
|
||||||
config: AppConfig
|
config: AppConfig
|
||||||
tool_context: ToolContext
|
tool_context: ToolContext
|
||||||
is_new: bool = True
|
|
||||||
state_key: str | None = None
|
state_key: str | None = None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
|
|
@ -74,57 +73,43 @@ class ChatDeps:
|
||||||
state_data = nested
|
state_data = nested
|
||||||
|
|
||||||
session_state = self.tool_context.get(SESSION_NAMESPACE, SessionState)
|
session_state = self.tool_context.get(SESSION_NAMESPACE, SessionState)
|
||||||
|
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
|
||||||
|
|
||||||
if self.is_new:
|
session_state.citations = [
|
||||||
# First request for this context: fully populate from client state
|
Citation(**c) if isinstance(c, dict) else c
|
||||||
if session_state is not None:
|
for c in state_data.get("citations", [])
|
||||||
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 = [
|
qa_session_state = self.tool_context.get(QA_SESSION_NAMESPACE, QASessionState)
|
||||||
Citation(**c) if isinstance(c, dict) else c
|
if qa_session_state is not None:
|
||||||
for c in state_data.get("citations", [])
|
if "qa_history" in state_data:
|
||||||
]
|
from haiku.rag.tools.qa import QAHistoryEntry
|
||||||
|
|
||||||
qa_session_state = self.tool_context.get(
|
qa_session_state.qa_history = [
|
||||||
QA_SESSION_NAMESPACE, QASessionState
|
QAHistoryEntry(**qa) if isinstance(qa, dict) else qa
|
||||||
)
|
for qa in state_data.get("qa_history", [])
|
||||||
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 = [
|
# Prefer server's session_context (background summarizer may
|
||||||
QAHistoryEntry(**qa) if isinstance(qa, dict) else qa
|
# have updated it since the client's last snapshot).
|
||||||
for qa in state_data.get("qa_history", [])
|
if not qa_session_state.session_context:
|
||||||
]
|
|
||||||
|
|
||||||
# Restore session_context from client
|
|
||||||
session_context = state_data.get("session_context")
|
session_context = state_data.get("session_context")
|
||||||
if isinstance(session_context, dict):
|
if isinstance(session_context, dict):
|
||||||
qa_session_state.session_context = SessionContext(
|
qa_session_state.session_context = SessionContext(
|
||||||
**session_context
|
**session_context
|
||||||
).summary
|
).summary
|
||||||
elif session_context is None:
|
|
||||||
qa_session_state.session_context = None
|
|
||||||
|
|
||||||
# Handle initial_context -> session_context for first message
|
# Handle initial_context -> session_context for first message
|
||||||
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 not qa_session_state.session_context:
|
if initial and not qa_session_state.session_context:
|
||||||
qa_session_state.session_context = initial
|
qa_session_state.session_context = initial
|
||||||
else:
|
|
||||||
# Returning request: only merge client-controlled fields
|
|
||||||
if session_state is not None:
|
|
||||||
if "document_filter" in state_data:
|
|
||||||
session_state.document_filter = state_data.get(
|
|
||||||
"document_filter", []
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def create_chat_agent(
|
def create_chat_agent(
|
||||||
|
|
|
||||||
|
|
@ -251,7 +251,6 @@ class ChatApp(App):
|
||||||
deps = ChatDeps(
|
deps = ChatDeps(
|
||||||
config=self.config,
|
config=self.config,
|
||||||
tool_context=self.tool_context,
|
tool_context=self.tool_context,
|
||||||
is_new=False,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
async with self.agent.run_stream(
|
async with self.agent.run_stream(
|
||||||
|
|
|
||||||
|
|
@ -48,6 +48,7 @@ __all__ = [
|
||||||
"QASessionState",
|
"QASessionState",
|
||||||
"QAHistoryEntry",
|
"QAHistoryEntry",
|
||||||
"create_qa_toolset",
|
"create_qa_toolset",
|
||||||
|
"run_qa_core",
|
||||||
"create_analysis_toolset",
|
"create_analysis_toolset",
|
||||||
"SESSION_NAMESPACE",
|
"SESSION_NAMESPACE",
|
||||||
"SessionState",
|
"SessionState",
|
||||||
|
|
|
||||||
|
|
@ -64,7 +64,6 @@ def test_chat_deps_initialization(temp_db_path):
|
||||||
|
|
||||||
assert deps.config is Config
|
assert deps.config is Config
|
||||||
assert deps.tool_context is context
|
assert deps.tool_context is context
|
||||||
assert deps.is_new is True
|
|
||||||
assert deps.state_key is None
|
assert deps.state_key is None
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -153,6 +152,42 @@ def test_chat_deps_state_setter_parses_session_context_dict():
|
||||||
assert qa_session_state.session_context == "Previous conversation summary"
|
assert qa_session_state.session_context == "Previous conversation summary"
|
||||||
|
|
||||||
|
|
||||||
|
def test_chat_deps_state_setter_preserves_server_session_context():
|
||||||
|
"""Test that server's session_context is preferred over client's stale value."""
|
||||||
|
from haiku.rag.tools.qa import QA_SESSION_NAMESPACE, QASessionState
|
||||||
|
|
||||||
|
context = ToolContext()
|
||||||
|
qa_state = QASessionState()
|
||||||
|
qa_state.session_context = "Fresh summary from background summarizer"
|
||||||
|
context.register(QA_SESSION_NAMESPACE, qa_state)
|
||||||
|
context.register(SESSION_NAMESPACE, SessionState())
|
||||||
|
|
||||||
|
deps = ChatDeps(config=Config, tool_context=context, state_key=AGUI_STATE_KEY)
|
||||||
|
|
||||||
|
# Client sends stale session_context
|
||||||
|
incoming_state = {
|
||||||
|
AGUI_STATE_KEY: {
|
||||||
|
"session_context": {
|
||||||
|
"summary": "Stale summary from client",
|
||||||
|
"last_updated": "2025-01-27T12:00:00",
|
||||||
|
},
|
||||||
|
"qa_history": [],
|
||||||
|
"citations": [],
|
||||||
|
"document_filter": [],
|
||||||
|
"citation_registry": {},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
deps.state = incoming_state
|
||||||
|
|
||||||
|
# Server's fresher session_context should be preserved
|
||||||
|
qa_session_state = context.get(QA_SESSION_NAMESPACE)
|
||||||
|
assert isinstance(qa_session_state, QASessionState)
|
||||||
|
assert (
|
||||||
|
qa_session_state.session_context == "Fresh summary from background summarizer"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_chat_session_state():
|
def test_chat_session_state():
|
||||||
"""Test ChatSessionState model."""
|
"""Test ChatSessionState model."""
|
||||||
state = ChatSessionState()
|
state = ChatSessionState()
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue