Fix coverage
This commit is contained in:
parent
f4735667df
commit
a870052ee9
2 changed files with 207 additions and 3 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in a new issue