beacon-trial-finder / tests /agents /test_research.py
KevinIsInCoding
fix: enforce patient's phase preference in API call, not LLM guidance
276f805
Raw
History Blame Contribute Delete
15.1 kB
"""Tests for agents/research.py.
run_research_agent: blocking loop — search tool routing, multi-call, error path.
stream_research_agent: FakeStream pattern — status/token/done event sequence.
search_trials_api, _flatten_and_rank, and bulk_parse_and_strip are patched so
tests are isolated to the loop/orchestration logic in research.py.
"""
from __future__ import annotations
from unittest.mock import MagicMock, patch
import pytest
from agents.research import run_research_agent, stream_research_agent, _phase_rank, _rank_and_slim
from models import PatientProfile
from tests.conftest import FakeStream, make_message, make_text_block, make_tool_use_block
# ---------------------------------------------------------------------------
# _phase_rank and _rank_and_slim
# ---------------------------------------------------------------------------
def _trial(study_type="INTERVENTIONAL", phase="PHASE3", distance=10.0, **extra):
return {"study_type": study_type, "phase": phase, "closest_site_miles": distance, **extra}
class TestPhaseRank:
def test_phase4_beats_phase3(self):
assert _phase_rank(_trial(phase="PHASE4")) < _phase_rank(_trial(phase="PHASE3"))
def test_phase3_beats_phase2(self):
assert _phase_rank(_trial(phase="PHASE3")) < _phase_rank(_trial(phase="PHASE2"))
def test_phase2_beats_phase1(self):
assert _phase_rank(_trial(phase="PHASE2")) < _phase_rank(_trial(phase="PHASE1"))
def test_phase1_beats_eap(self):
assert _phase_rank(_trial(phase="PHASE1")) < _phase_rank(_trial(study_type="EXPANDED_ACCESS", phase=""))
def test_eap_beats_observational(self):
assert _phase_rank(_trial(study_type="EXPANDED_ACCESS", phase="")) < _phase_rank(_trial(study_type="OBSERVATIONAL", phase=""))
def test_na_phase_interventional_between_phase1_and_eap(self):
rank_na = _phase_rank(_trial(phase="NA"))
assert _phase_rank(_trial(phase="PHASE1")) < rank_na
assert rank_na < _phase_rank(_trial(study_type="EXPANDED_ACCESS", phase=""))
class TestRankAndSlim:
def test_phase4_sorted_before_phase3(self):
trials = [_trial(phase="PHASE3", nct_id="B"), _trial(phase="PHASE4", nct_id="A")]
result = _rank_and_slim(trials)
assert result[0]["nct_id"] == "A"
def test_within_same_phase_closer_first(self):
trials = [_trial(phase="PHASE3", distance=50.0, nct_id="far"), _trial(phase="PHASE3", distance=5.0, nct_id="near")]
result = _rank_and_slim(trials)
assert result[0]["nct_id"] == "near"
def test_capped_at_max_trials(self):
trials = [_trial(phase="PHASE2", nct_id=str(i)) for i in range(20)]
result = _rank_and_slim(trials)
assert len(result) <= 15
def test_strips_summary_and_conditions(self):
trials = [_trial(phase="PHASE3", summary="long text", conditions=["ALS"])]
result = _rank_and_slim(trials)
assert "summary" not in result[0]
assert "conditions" not in result[0]
def test_nearest_sites_capped_at_3(self):
sites = [{"label": f"Site {i}"} for i in range(5)]
trials = [_trial(phase="PHASE3", nearest_sites=sites)]
result = _rank_and_slim(trials)
assert len(result[0]["nearest_sites"]) == 3
def test_intervention_description_stripped(self):
iv = [{"type": "DRUG", "name": "DrugX", "description": "long description text"}]
trials = [_trial(phase="PHASE3", interventions=iv)]
result = _rank_and_slim(trials)
assert "description" not in result[0]["interventions"][0]
assert result[0]["interventions"][0]["name"] == "DrugX"
PATCH_SEARCH = "agents.research.search_trials_api"
PATCH_FLATTEN = "agents.research._flatten_and_rank"
PATCH_BULK = "agents.research.bulk_parse_and_strip"
def _search_block(condition: str = "ALS") -> MagicMock:
return make_tool_use_block(
"search_clinical_trials",
{"condition": condition, "lat": 42.33, "lon": -71.10, "radius_miles": 100},
tool_use_id="tu_search1",
)
# ---------------------------------------------------------------------------
# run_research_agent
# ---------------------------------------------------------------------------
class TestRunResearchAgent:
def test_returns_text_on_end_turn(self, mock_client, als_patient):
final_msg = make_message(
content=[make_text_block("Here are your trials.")],
stop_reason="end_turn",
)
mock_client.messages.create.return_value = final_msg
result = run_research_agent(mock_client, als_patient)
assert result == "Here are your trials."
def test_no_trials_text_fallback(self, mock_client, als_patient):
final_msg = make_message(content=[], stop_reason="end_turn")
mock_client.messages.create.return_value = final_msg
result = run_research_agent(mock_client, als_patient)
assert result == "No analysis produced."
def test_one_search_call_then_end_turn(self, mock_client, als_patient):
search_msg = make_message(
content=[_search_block()], stop_reason="tool_use"
)
final_msg = make_message(
content=[make_text_block("Found 3 trials.")], stop_reason="end_turn"
)
mock_client.messages.create.side_effect = [search_msg, final_msg]
with patch(PATCH_SEARCH, return_value=[]) as ms, \
patch(PATCH_FLATTEN, return_value=[]) as mf, \
patch(PATCH_BULK, return_value=[]) as mb:
result = run_research_agent(mock_client, als_patient)
assert result == "Found 3 trials."
ms.assert_called_once()
mf.assert_called_once()
mb.assert_called_once()
def test_two_search_calls_before_end_turn(self, mock_client, als_patient):
search_msg1 = make_message(content=[_search_block("ALS")], stop_reason="tool_use")
search_msg2 = make_message(content=[_search_block("Motor Neuron Disease")], stop_reason="tool_use")
final_msg = make_message(
content=[make_text_block("Done.")], stop_reason="end_turn"
)
mock_client.messages.create.side_effect = [search_msg1, search_msg2, final_msg]
with patch(PATCH_SEARCH, return_value=[]), \
patch(PATCH_FLATTEN, return_value=[]), \
patch(PATCH_BULK, return_value=[]):
result = run_research_agent(mock_client, als_patient)
assert result == "Done."
def test_profile_phases_enforced_in_interventional_search(self, mock_client, als_patient):
"""Patient's phase preference is passed to search_trials_api, not whatever LLM chose."""
patient_with_phases = als_patient.__class__(
**{**als_patient.__dict__, "phases": ["3", "4"]}
)
# LLM passes no phases in its tool call args
search_msg = make_message(content=[_search_block()], stop_reason="tool_use")
final_msg = make_message(content=[make_text_block("Done.")], stop_reason="end_turn")
mock_client.messages.create.side_effect = [search_msg, final_msg]
with patch(PATCH_SEARCH, return_value=[]) as ms, \
patch(PATCH_FLATTEN, return_value=[]), \
patch(PATCH_BULK, return_value=[]):
run_research_agent(mock_client, patient_with_phases)
_, kwargs = ms.call_args
assert kwargs["phases"] == ["3", "4"]
def test_no_phase_filter_when_profile_phases_empty(self, mock_client, als_patient):
"""No phase filter applied when patient has no phase preference."""
patient_no_phases = als_patient.__class__(
**{**als_patient.__dict__, "phases": []}
)
search_msg = make_message(content=[_search_block()], stop_reason="tool_use")
final_msg = make_message(content=[make_text_block("Done.")], stop_reason="end_turn")
mock_client.messages.create.side_effect = [search_msg, final_msg]
with patch(PATCH_SEARCH, return_value=[]) as ms, \
patch(PATCH_FLATTEN, return_value=[]), \
patch(PATCH_BULK, return_value=[]):
run_research_agent(mock_client, patient_no_phases)
_, kwargs = ms.call_args
assert kwargs["phases"] is None
def test_search_results_serialized_into_tool_result(self, mock_client, als_patient):
"""Verify flatten output flows through bulk_parse and into the tool-result message."""
import json as _json
fake_trial = {"nct_id": "NCT00000001", "title": "ALS Trial", "parsed_criteria": []}
search_msg = make_message(content=[_search_block()], stop_reason="tool_use")
final_msg = make_message(
content=[make_text_block("Here are results.")], stop_reason="end_turn"
)
mock_client.messages.create.side_effect = [search_msg, final_msg]
with patch(PATCH_SEARCH, return_value=[{}]), \
patch(PATCH_FLATTEN, return_value=[fake_trial]), \
patch(PATCH_BULK, return_value=[fake_trial]):
run_research_agent(mock_client, als_patient)
second_call_msgs = mock_client.messages.create.call_args_list[1][1]["messages"]
user_msgs = [m for m in second_call_msgs if m["role"] == "user"]
tool_blocks = [
b for m in user_msgs
for b in (m["content"] if isinstance(m["content"], list) else [])
if isinstance(b, dict) and b.get("type") == "tool_result"
]
assert tool_blocks, "No tool_result block found in second call"
content = _json.loads(tool_blocks[0]["content"])
assert content[0]["nct_id"] == "NCT00000001"
def test_search_api_error_passes_error_content(self, mock_client, als_patient):
search_msg = make_message(content=[_search_block()], stop_reason="tool_use")
final_msg = make_message(
content=[make_text_block("API was unavailable.")], stop_reason="end_turn"
)
mock_client.messages.create.side_effect = [search_msg, final_msg]
with patch(PATCH_SEARCH, side_effect=Exception("timeout")), \
patch(PATCH_FLATTEN, return_value=[]), \
patch(PATCH_BULK, return_value=[]):
result = run_research_agent(mock_client, als_patient)
assert result == "API was unavailable."
# messages is passed by reference and mutated after each call, so search
# the captured list for the tool_result block rather than relying on index
second_call_msgs = mock_client.messages.create.call_args_list[1][1]["messages"]
user_msgs = [m for m in second_call_msgs if m["role"] == "user"]
tool_result_blocks = [
b for m in user_msgs
for b in (m["content"] if isinstance(m["content"], list) else [])
if isinstance(b, dict) and b.get("type") == "tool_result"
]
assert any(b.get("is_error") is True for b in tool_result_blocks)
# ---------------------------------------------------------------------------
# stream_research_agent
# ---------------------------------------------------------------------------
class TestStreamResearchAgent:
def test_yields_token_events(self, mock_client, als_patient):
final_msg = make_message(
content=[make_text_block("Here are 3 trials.")],
stop_reason="end_turn",
)
mock_client.messages.stream.return_value = FakeStream(
tokens=["Here ", "are ", "3 trials."],
final_message=final_msg,
)
events = list(stream_research_agent(mock_client, als_patient))
token_events = [e for e in events if e[0] == "token"]
assert token_events == [("token", "Here "), ("token", "are "), ("token", "3 trials.")]
def test_yields_done_event_on_end_turn(self, mock_client, als_patient):
final_msg = make_message(
content=[make_text_block("Analysis complete.")],
stop_reason="end_turn",
)
mock_client.messages.stream.return_value = FakeStream(
tokens=["Analysis complete."],
final_message=final_msg,
)
events = list(stream_research_agent(mock_client, als_patient))
done_events = [e for e in events if e[0] == "done"]
assert len(done_events) == 1
assert done_events[0][1] == "Analysis complete."
def test_yields_status_event_before_search(self, mock_client, als_patient):
search_msg = make_message(content=[_search_block()], stop_reason="tool_use")
final_msg = make_message(
content=[make_text_block("Done.")], stop_reason="end_turn"
)
mock_client.messages.stream.side_effect = [
FakeStream(tokens=[], final_message=search_msg),
FakeStream(tokens=["Done."], final_message=final_msg),
]
with patch(PATCH_SEARCH, return_value=[]), \
patch(PATCH_FLATTEN, return_value=[]), \
patch(PATCH_BULK, return_value=[]):
events = list(stream_research_agent(mock_client, als_patient))
status_events = [e for e in events if e[0] == "status"]
assert len(status_events) >= 1
def test_terminates_with_done_as_last_event(self, mock_client, als_patient):
final_msg = make_message(
content=[make_text_block("Finished.")], stop_reason="end_turn"
)
mock_client.messages.stream.return_value = FakeStream(
tokens=["Finished."], final_message=final_msg
)
events = list(stream_research_agent(mock_client, als_patient))
assert events[-1][0] == "done"
def test_search_api_error_still_yields_done(self, mock_client, als_patient):
search_msg = make_message(content=[_search_block()], stop_reason="tool_use")
final_msg = make_message(
content=[make_text_block("API was down.")], stop_reason="end_turn"
)
mock_client.messages.stream.side_effect = [
FakeStream(tokens=[], final_message=search_msg),
FakeStream(tokens=["API was down."], final_message=final_msg),
]
with patch(PATCH_SEARCH, side_effect=Exception("connect timeout")), \
patch(PATCH_FLATTEN, return_value=[]), \
patch(PATCH_BULK, return_value=[]):
events = list(stream_research_agent(mock_client, als_patient))
assert events[-1][0] == "done"
def test_two_stream_calls_for_one_search_cycle(self, mock_client, als_patient):
search_msg = make_message(content=[_search_block()], stop_reason="tool_use")
final_msg = make_message(
content=[make_text_block("Results.")], stop_reason="end_turn"
)
mock_client.messages.stream.side_effect = [
FakeStream(tokens=[], final_message=search_msg),
FakeStream(tokens=["Results."], final_message=final_msg),
]
with patch(PATCH_SEARCH, return_value=[]), \
patch(PATCH_FLATTEN, return_value=[]), \
patch(PATCH_BULK, return_value=[]):
list(stream_research_agent(mock_client, als_patient))
assert mock_client.messages.stream.call_count == 2