Fix session_id not persisting across AG-UI requests
This commit is contained in:
parent
a889883cfc
commit
d1a9439500
3 changed files with 35 additions and 21 deletions
|
|
@ -1,4 +1,3 @@
|
||||||
import uuid
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
@ -75,7 +74,7 @@ class SessionContext(BaseModel):
|
||||||
class ChatSessionState(BaseModel):
|
class ChatSessionState(BaseModel):
|
||||||
"""State shared between frontend and agent via AG-UI."""
|
"""State shared between frontend and agent via AG-UI."""
|
||||||
|
|
||||||
session_id: str = Field(default_factory=lambda: str(uuid.uuid4()))
|
session_id: str = ""
|
||||||
initial_context: str | None = None
|
initial_context: str | None = None
|
||||||
citations: list[Citation] = []
|
citations: list[Citation] = []
|
||||||
qa_history: list[QAResponse] = []
|
qa_history: list[QAResponse] = []
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,3 @@
|
||||||
import uuid
|
|
||||||
|
|
||||||
from ag_ui.core import StateDeltaEvent
|
from ag_ui.core import StateDeltaEvent
|
||||||
|
|
||||||
from haiku.rag.agents.chat.state import (
|
from haiku.rag.agents.chat.state import (
|
||||||
|
|
@ -578,13 +576,14 @@ def test_chat_deps_state_getter_includes_document_filter():
|
||||||
assert state[AGUI_STATE_KEY]["document_filter"] == ["doc1.pdf", "doc2.pdf"]
|
assert state[AGUI_STATE_KEY]["document_filter"] == ["doc1.pdf", "doc2.pdf"]
|
||||||
|
|
||||||
|
|
||||||
def test_chat_session_state_auto_generates_session_id():
|
def test_chat_session_state_defaults_to_empty_session_id():
|
||||||
"""New ChatSessionState should have a valid UUID session_id."""
|
"""New ChatSessionState should default to empty session_id.
|
||||||
|
|
||||||
|
Tools in agent.py detect the empty string and assign a UUID,
|
||||||
|
which then appears in the state delta so clients receive it.
|
||||||
|
"""
|
||||||
state = ChatSessionState()
|
state = ChatSessionState()
|
||||||
assert state.session_id
|
assert state.session_id == ""
|
||||||
assert len(state.session_id) == 36 # UUID format
|
|
||||||
# Verify it's a valid UUID
|
|
||||||
uuid.UUID(state.session_id)
|
|
||||||
|
|
||||||
|
|
||||||
def test_chat_session_state_preserves_explicit_session_id():
|
def test_chat_session_state_preserves_explicit_session_id():
|
||||||
|
|
@ -593,13 +592,6 @@ def test_chat_session_state_preserves_explicit_session_id():
|
||||||
assert state.session_id == "my-custom-id"
|
assert state.session_id == "my-custom-id"
|
||||||
|
|
||||||
|
|
||||||
def test_chat_session_state_each_instance_gets_unique_id():
|
|
||||||
"""Each new instance should get a unique session_id."""
|
|
||||||
state1 = ChatSessionState()
|
|
||||||
state2 = ChatSessionState()
|
|
||||||
assert state1.session_id != state2.session_id
|
|
||||||
|
|
||||||
|
|
||||||
def test_chat_session_state_initial_context_default_none():
|
def test_chat_session_state_initial_context_default_none():
|
||||||
"""Initial context should default to None."""
|
"""Initial context should default to None."""
|
||||||
state = ChatSessionState()
|
state = ChatSessionState()
|
||||||
|
|
@ -650,6 +642,27 @@ def test_chat_session_state_model_dump_json_serializes_datetime():
|
||||||
assert snapshot["session_context"]["last_updated"] == "2025-01-27T12:00:00"
|
assert snapshot["session_context"]["last_updated"] == "2025-01-27T12:00:00"
|
||||||
|
|
||||||
|
|
||||||
|
def test_emit_state_event_includes_session_id_when_assigned():
|
||||||
|
"""emit_state_event detects session_id change from empty to UUID.
|
||||||
|
|
||||||
|
When session_id defaults to "" and the tool assigns a UUID,
|
||||||
|
the delta must include session_id so clients can persist it.
|
||||||
|
"""
|
||||||
|
from haiku.rag.agents.chat.state import emit_state_event
|
||||||
|
|
||||||
|
current_state = ChatSessionState() # session_id=""
|
||||||
|
new_state = ChatSessionState(session_id="assigned-uuid-123")
|
||||||
|
|
||||||
|
event = emit_state_event(current_state, new_state)
|
||||||
|
|
||||||
|
assert event is not None
|
||||||
|
session_id_op = next(
|
||||||
|
(op for op in event.delta if op["path"] == "/session_id"), None
|
||||||
|
)
|
||||||
|
assert session_id_op is not None
|
||||||
|
assert session_id_op["value"] == "assigned-uuid-123"
|
||||||
|
|
||||||
|
|
||||||
def test_emit_state_event_returns_none_when_no_changes():
|
def test_emit_state_event_returns_none_when_no_changes():
|
||||||
"""emit_state_event returns None when states are identical."""
|
"""emit_state_event returns None when states are identical."""
|
||||||
from haiku.rag.agents.chat.state import emit_state_event
|
from haiku.rag.agents.chat.state import emit_state_event
|
||||||
|
|
|
||||||
|
|
@ -266,9 +266,9 @@ async def test_clear_chat_resets_session(temp_db_path: Path):
|
||||||
await chat_history.add_message("assistant", "Hi there")
|
await chat_history.add_message("assistant", "Hi there")
|
||||||
assert len(chat_history.messages) == 2
|
assert len(chat_history.messages) == 2
|
||||||
|
|
||||||
# Store original session ID
|
# Mutate session state to simulate an active session
|
||||||
assert app.session_state is not None
|
assert app.session_state is not None
|
||||||
original_session_id = app.session_state.session_id
|
app.session_state.session_id = "active-session-123"
|
||||||
|
|
||||||
# Clear chat via action (available through command palette)
|
# Clear chat via action (available through command palette)
|
||||||
await app.action_clear_chat()
|
await app.action_clear_chat()
|
||||||
|
|
@ -277,9 +277,11 @@ async def test_clear_chat_resets_session(temp_db_path: Path):
|
||||||
# Verify messages cleared
|
# Verify messages cleared
|
||||||
assert len(chat_history.messages) == 0
|
assert len(chat_history.messages) == 0
|
||||||
|
|
||||||
# Verify session state reset (new session ID)
|
# Verify session state reset
|
||||||
assert app.session_state is not None
|
assert app.session_state is not None
|
||||||
assert app.session_state.session_id != original_session_id
|
assert app.session_state.session_id == ""
|
||||||
|
assert app.session_state.qa_history == []
|
||||||
|
assert app.session_state.citations == []
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue