diff --git a/haiku_rag_slim/haiku/rag/tools/qa.py b/haiku_rag_slim/haiku/rag/tools/qa.py index 2f9d67c8..76d301ac 100644 --- a/haiku_rag_slim/haiku/rag/tools/qa.py +++ b/haiku_rag_slim/haiku/rag/tools/qa.py @@ -267,9 +267,7 @@ def create_qa_toolset( if client_snapshot is not None and tool_context is not None: new_snapshot = tool_context.build_state_snapshot() state_event = compute_combined_state_delta( - client_snapshot, - new_snapshot, - state_key=state_key, + client_snapshot, new_snapshot, state_key=state_key ) if state_event is not None: diff --git a/tests/agents/chat/test_chat_agent.py b/tests/agents/chat/test_chat_agent.py index 3a36d80a..659ad1ff 100644 --- a/tests/agents/chat/test_chat_agent.py +++ b/tests/agents/chat/test_chat_agent.py @@ -1419,3 +1419,209 @@ async def test_summarization_task_cleanup_on_completion(): # Task should be cleaned up assert key not in _summarization_tasks + + +@pytest.mark.asyncio +async def test_ask_tool_returns_qa_result_when_no_state_delta(temp_db_path): + """Test ask tool falls through to QAResult when state delta is empty. + + When run_qa_core produces no state changes (client_snapshot == new_snapshot), + compute_combined_state_delta returns None and the tool returns QAResult directly. + """ + from unittest.mock import AsyncMock, patch + + from pydantic_ai import Agent + from pydantic_ai.models.test import TestModel + + from haiku.rag.tools.models import QAResult + from haiku.rag.tools.qa import create_qa_toolset + + async with HaikuRAG(temp_db_path, create=True) as client: + context = ToolContext() + prepare_chat_context(context, features=["search", "qa"]) + context.state_key = AGUI_STATE_KEY + + deps = ChatDeps(config=Config, client=client, tool_context=context) + toolset = create_qa_toolset(Config) + agent = Agent( + TestModel(call_tools=["ask"]), + deps_type=ChatDeps, + toolsets=[toolset], # ty: ignore[invalid-argument-type] + ) + + # Mock run_qa_core to return a result WITHOUT modifying any state. + # Since client_snapshot == new_snapshot, the delta is None. + mock_result = QAResult( + question="test", answer="The answer", confidence=0.9, citations=[] + ) + with patch( + "haiku.rag.tools.qa.run_qa_core", + new_callable=AsyncMock, + return_value=mock_result, + ): + result = await agent.run( + "test question", + deps=deps, # ty: ignore[invalid-argument-type] + ) + + # The tool should have returned the QAResult (not a ToolReturn) + assert "The answer" in result.output + + +@pytest.mark.asyncio +async def test_ask_tool_delta_without_citations(temp_db_path): + """Test ask tool emits StateDeltaEvent without citation sources when citations are empty. + + When the QA result has no citations but state DID change, + the tool returns a ToolReturn with a delta but no 'Sources:' line. + """ + from unittest.mock import AsyncMock, patch + + from pydantic_ai import Agent + from pydantic_ai.models.test import TestModel + + from haiku.rag.tools.models import QAResult + from haiku.rag.tools.qa import ( + QA_SESSION_NAMESPACE, + QASessionState, + create_qa_toolset, + ) + + async with HaikuRAG(temp_db_path, create=True) as client: + context = ToolContext() + prepare_chat_context(context, features=["search", "qa"]) + context.state_key = AGUI_STATE_KEY + + deps = ChatDeps(config=Config, client=client, tool_context=context) + toolset = create_qa_toolset(Config) + agent = Agent( + TestModel(call_tools=["ask"]), + deps_type=ChatDeps, + toolsets=[toolset], # ty: ignore[invalid-argument-type] + ) + + mock_result = QAResult( + question="test", answer="No citations answer", confidence=0.9, citations=[] + ) + + async def mock_qa_core(*args, **kwargs): + # Simulate state mutation (qa_history entry) but no citations + ctx = kwargs.get("context") + if ctx is not None: + qa_state = ctx.get(QA_SESSION_NAMESPACE, QASessionState) + if qa_state is not None: + qa_state.qa_history.append( + QAHistoryEntry(question="test", answer="No citations answer") + ) + return mock_result + + with patch( + "haiku.rag.tools.qa.run_qa_core", + new_callable=AsyncMock, + side_effect=mock_qa_core, + ): + result = await agent.run( + "test question", + deps=deps, # ty: ignore[invalid-argument-type] + ) + + # Should have the answer but NO "Sources:" line + assert "No citations answer" in result.output + assert "Sources:" not in result.output + + # Should have emitted a StateDeltaEvent (qa_history changed) + state = extract_state_from_result(result) + assert state is not None + assert len(state.get("qa_history", [])) > 0 + + +@pytest.mark.asyncio +async def test_ask_tool_delta_with_citations(temp_db_path): + """Test ask tool emits StateDeltaEvent with 'Sources:' line when citations are present. + + When the QA result has citations and state changed, the tool returns + a ToolReturn with a delta and appended 'Sources: [1] [2]' line. + """ + from unittest.mock import AsyncMock, patch + + from pydantic_ai import Agent + from pydantic_ai.models.test import TestModel + + from haiku.rag.tools.models import QAResult + from haiku.rag.tools.qa import ( + QA_SESSION_NAMESPACE, + QASessionState, + create_qa_toolset, + ) + + async with HaikuRAG(temp_db_path, create=True) as client: + context = ToolContext() + prepare_chat_context(context, features=["search", "qa"]) + context.state_key = AGUI_STATE_KEY + + deps = ChatDeps(config=Config, client=client, tool_context=context) + toolset = create_qa_toolset(Config) + agent = Agent( + TestModel(call_tools=["ask"]), + deps_type=ChatDeps, + toolsets=[toolset], # ty: ignore[invalid-argument-type] + ) + + citations = [ + Citation( + index=1, + document_id="doc-1", + chunk_id="chunk-a", + document_uri="test.md", + document_title="Test", + content="content", + ), + Citation( + index=2, + document_id="doc-1", + chunk_id="chunk-b", + document_uri="test.md", + document_title="Test", + content="content", + ), + ] + mock_result = QAResult( + question="test", + answer="Cited answer", + confidence=0.95, + citations=citations, + ) + + async def mock_qa_core(*args, **kwargs): + ctx = kwargs.get("context") + if ctx is not None: + qa_state = ctx.get(QA_SESSION_NAMESPACE, QASessionState) + if qa_state is not None: + qa_state.qa_history.append( + QAHistoryEntry( + question="test", + answer="Cited answer", + citations=citations, + ) + ) + return mock_result + + with patch( + "haiku.rag.tools.qa.run_qa_core", + new_callable=AsyncMock, + side_effect=mock_qa_core, + ): + result = await agent.run( + "test question", + deps=deps, # ty: ignore[invalid-argument-type] + ) + + # Should have the answer WITH "Sources:" line + assert "Cited answer" in result.output + assert "Sources:" in result.output + assert "[1]" in result.output + assert "[2]" in result.output + + # Should have emitted a StateDeltaEvent + state = extract_state_from_result(result) + assert state is not None