Add document_filter to ChatSessionState for session-level filtering
This commit is contained in:
parent
0f0c3a076f
commit
5b14038dc5
2 changed files with 136 additions and 0 deletions
|
|
@ -48,6 +48,7 @@ class ChatSessionState(BaseModel):
|
|||
citations: list[Citation] = []
|
||||
qa_history: list[QAResponse] = []
|
||||
session_context: SessionContext | None = None
|
||||
document_filter: list[str] = []
|
||||
|
||||
|
||||
@dataclass
|
||||
|
|
@ -98,6 +99,10 @@ class ChatDeps:
|
|||
]
|
||||
if state_data.get("session_id"):
|
||||
self.session_state.session_id = state_data["session_id"]
|
||||
if "document_filter" in state_data:
|
||||
self.session_state.document_filter = state_data.get(
|
||||
"document_filter", []
|
||||
)
|
||||
# NOTE: session_context intentionally NOT updated from client
|
||||
# The agent owns this via server-side cache
|
||||
|
||||
|
|
@ -120,3 +125,23 @@ def build_document_filter(document_name: str) -> str:
|
|||
f"LOWER(uri) LIKE LOWER('%{escaped}%') OR LOWER(title) LIKE LOWER('%{escaped}%') "
|
||||
f"OR LOWER(uri) LIKE LOWER('%{no_spaces}%') OR LOWER(title) LIKE LOWER('%{no_spaces}%')"
|
||||
)
|
||||
|
||||
|
||||
def build_multi_document_filter(document_names: list[str]) -> str | None:
|
||||
"""Build SQL filter for multiple document names (OR combined)."""
|
||||
if not document_names:
|
||||
return None
|
||||
filters = [build_document_filter(name) for name in document_names]
|
||||
if len(filters) == 1:
|
||||
return filters[0]
|
||||
return " OR ".join(f"({f})" for f in filters)
|
||||
|
||||
|
||||
def combine_filters(filter1: str | None, filter2: str | None) -> str | None:
|
||||
"""Combine two SQL filters with AND logic."""
|
||||
filters = [f for f in [filter1, filter2] if f]
|
||||
if not filters:
|
||||
return None
|
||||
if len(filters) == 1:
|
||||
return filters[0]
|
||||
return f"({filters[0]}) AND ({filters[1]})"
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@ from haiku.rag.agents.chat.state import (
|
|||
MAX_QA_HISTORY,
|
||||
QAResponse,
|
||||
build_document_filter,
|
||||
build_multi_document_filter,
|
||||
combine_filters,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -29,6 +31,56 @@ def test_build_document_filter_escapes_quotes():
|
|||
assert "O''Reilly" in result
|
||||
|
||||
|
||||
def test_build_multi_document_filter_empty():
|
||||
"""Test build_multi_document_filter returns None for empty list."""
|
||||
result = build_multi_document_filter([])
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_build_multi_document_filter_single():
|
||||
"""Test build_multi_document_filter with single document."""
|
||||
result = build_multi_document_filter(["mytest"])
|
||||
assert result is not None
|
||||
assert "LOWER(uri) LIKE LOWER('%mytest%')" in result
|
||||
assert "LOWER(title) LIKE LOWER('%mytest%')" in result
|
||||
# Single document should not have extra wrapping parentheses
|
||||
assert " OR (" not in result
|
||||
|
||||
|
||||
def test_build_multi_document_filter_multiple():
|
||||
"""Test build_multi_document_filter with multiple documents."""
|
||||
result = build_multi_document_filter(["doc1", "doc2"])
|
||||
assert result is not None
|
||||
# Should have OR-combined filters
|
||||
assert "doc1" in result
|
||||
assert "doc2" in result
|
||||
assert " OR (" in result
|
||||
|
||||
|
||||
def test_combine_filters_both_none():
|
||||
"""Test combine_filters with both None."""
|
||||
result = combine_filters(None, None)
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_combine_filters_first_only():
|
||||
"""Test combine_filters with only first filter."""
|
||||
result = combine_filters("uri = 'test'", None)
|
||||
assert result == "uri = 'test'"
|
||||
|
||||
|
||||
def test_combine_filters_second_only():
|
||||
"""Test combine_filters with only second filter."""
|
||||
result = combine_filters(None, "title = 'doc'")
|
||||
assert result == "title = 'doc'"
|
||||
|
||||
|
||||
def test_combine_filters_both():
|
||||
"""Test combine_filters combines with AND."""
|
||||
result = combine_filters("uri = 'test'", "title = 'doc'")
|
||||
assert result == "(uri = 'test') AND (title = 'doc')"
|
||||
|
||||
|
||||
def test_max_qa_history_constant():
|
||||
"""Test MAX_QA_HISTORY constant value."""
|
||||
assert MAX_QA_HISTORY == 50
|
||||
|
|
@ -322,3 +374,62 @@ def test_chat_deps_state_setter_ignores_session_context():
|
|||
assert deps.session_state is not None
|
||||
assert deps.session_state.session_context is not None
|
||||
assert deps.session_state.session_context.summary == "Server-side context"
|
||||
|
||||
|
||||
def test_chat_deps_state_setter_restores_document_filter():
|
||||
"""Test ChatDeps.state setter restores document_filter from incoming state."""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from haiku.rag.agents.chat.state import AGUI_STATE_KEY, ChatDeps, ChatSessionState
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_config = MagicMock()
|
||||
|
||||
session_state = ChatSessionState(session_id="test")
|
||||
deps = ChatDeps(
|
||||
client=mock_client,
|
||||
config=mock_config,
|
||||
session_state=session_state,
|
||||
state_key=AGUI_STATE_KEY,
|
||||
)
|
||||
|
||||
incoming_state = {
|
||||
AGUI_STATE_KEY: {
|
||||
"session_id": "test",
|
||||
"qa_history": [],
|
||||
"citations": [],
|
||||
"document_filter": ["doc1.pdf", "doc2.pdf"],
|
||||
}
|
||||
}
|
||||
|
||||
deps.state = incoming_state
|
||||
|
||||
assert deps.session_state is not None
|
||||
assert deps.session_state.document_filter == ["doc1.pdf", "doc2.pdf"]
|
||||
|
||||
|
||||
def test_chat_deps_state_getter_includes_document_filter():
|
||||
"""Test ChatDeps.state getter includes document_filter."""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from haiku.rag.agents.chat.state import AGUI_STATE_KEY, ChatDeps, ChatSessionState
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_config = MagicMock()
|
||||
|
||||
session_state = ChatSessionState(
|
||||
session_id="test-123",
|
||||
document_filter=["doc1.pdf", "doc2.pdf"],
|
||||
)
|
||||
|
||||
deps = ChatDeps(
|
||||
client=mock_client,
|
||||
config=mock_config,
|
||||
session_state=session_state,
|
||||
state_key=AGUI_STATE_KEY,
|
||||
)
|
||||
|
||||
state = deps.state
|
||||
assert state is not None
|
||||
assert AGUI_STATE_KEY in state
|
||||
assert state[AGUI_STATE_KEY]["document_filter"] == ["doc1.pdf", "doc2.pdf"]
|
||||
|
|
|
|||
Loading…
Reference in a new issue