beacon-trial-finder / tests /test_trials_api.py
KevinIsInCoding
test(critic-fixes): address brittleness, over-mocking, and missing edge cases (#20)
3d5e588 unverified
Raw
History Blame Contribute Delete
12.4 kB
"""Tests for trials_api.py.
_flatten_and_rank: pure dict reshaping — no mocks.
search_trials_api: HTTP is mocked with pytest-httpx.
"""
from __future__ import annotations
import json
import httpx
import pytest
from trials_api import _flatten_and_rank, search_trials_api
# ---------------------------------------------------------------------------
# Minimal study builder
# ---------------------------------------------------------------------------
def make_study(
nct_id: str = "NCT00000001",
title: str = "Test Trial",
phase: list[str] | None = None,
sponsor: str = "ACME Pharma",
locations: list[dict] | None = None,
eligibility_criteria: str = "Inclusion:\n- Age 18-75",
) -> dict:
return {
"protocolSection": {
"identificationModule": {"nctId": nct_id, "briefTitle": title},
"descriptionModule": {"briefSummary": "A test trial."},
"eligibilityModule": {
"eligibilityCriteria": eligibility_criteria,
"minimumAge": "18 Years",
"maximumAge": "75 Years",
"sex": "ALL",
"healthyVolunteers": "No",
"stdAges": ["ADULT"],
},
"contactsLocationsModule": {
"centralContacts": [{"phone": "617-555-0100", "email": "pi@hospital.org"}],
"overallOfficials": [{"name": "Dr. Smith", "role": "PRINCIPAL_INVESTIGATOR"}],
"locations": locations or [],
},
"sponsorCollaboratorsModule": {"leadSponsor": {"name": sponsor}},
"designModule": {
"studyType": "INTERVENTIONAL",
"phases": phase or ["PHASE2"],
"enrollmentInfo": {"count": 50},
},
"conditionsModule": {"conditions": ["ALS"], "keywords": ["motor neuron"]},
"armsInterventionsModule": {"interventions": []},
}
}
def make_location(lat: float, lon: float, city: str = "Boston", state: str = "MA") -> dict:
return {
"facility": f"{city} Medical Center",
"city": city,
"state": state,
"geoPoint": {"lat": lat, "lon": lon},
"contacts": [{"phone": "617-555-0200", "email": "site@hospital.org"}],
}
# ---------------------------------------------------------------------------
# _flatten_and_rank — pure logic, no HTTP
# ---------------------------------------------------------------------------
class TestFlattenAndRank:
PATIENT_LAT = 42.3370
PATIENT_LON = -71.1061 # Boston
def test_empty_input_returns_empty(self):
assert _flatten_and_rank([], self.PATIENT_LAT, self.PATIENT_LON) == []
def test_nct_id_extracted(self):
studies = [make_study(nct_id="NCT12345678")]
result = _flatten_and_rank(studies, self.PATIENT_LAT, self.PATIENT_LON)
assert result[0]["nct_id"] == "NCT12345678"
def test_title_extracted(self):
studies = [make_study(title="My Trial")]
result = _flatten_and_rank(studies, self.PATIENT_LAT, self.PATIENT_LON)
assert result[0]["title"] == "My Trial"
def test_sponsor_extracted(self):
studies = [make_study(sponsor="Test Corp")]
result = _flatten_and_rank(studies, self.PATIENT_LAT, self.PATIENT_LON)
assert result[0]["sponsor"] == "Test Corp"
def test_no_locations_yields_none_closest(self):
studies = [make_study()]
result = _flatten_and_rank(studies, self.PATIENT_LAT, self.PATIENT_LON)
assert result[0]["closest_site_miles"] is None
def test_location_distance_computed(self):
loc = make_location(lat=42.3370, lon=-71.1061) # same as patient
studies = [make_study(locations=[loc])]
result = _flatten_and_rank(studies, self.PATIENT_LAT, self.PATIENT_LON)
assert result[0]["closest_site_miles"] == 0.0
def test_sorted_by_closest_distance(self):
near_loc = make_location(lat=42.3370, lon=-71.1061, city="Boston") # 0 mi
far_loc = make_location(lat=40.7128, lon=-74.0060, city="NewYork") # ~190 mi
near_study = make_study(nct_id="NCT00000001", locations=[near_loc])
far_study = make_study(nct_id="NCT00000002", locations=[far_loc])
# Pass far first to confirm sort works
result = _flatten_and_rank([far_study, near_study], self.PATIENT_LAT, self.PATIENT_LON)
assert result[0]["nct_id"] == "NCT00000001"
assert result[1]["nct_id"] == "NCT00000002"
def test_no_location_sorts_last(self):
loc = make_location(lat=42.3370, lon=-71.1061)
study_with = make_study(nct_id="NCT00000001", locations=[loc])
study_without = make_study(nct_id="NCT00000002", locations=[])
result = _flatten_and_rank([study_without, study_with], self.PATIENT_LAT, self.PATIENT_LON)
assert result[0]["nct_id"] == "NCT00000001"
assert result[1]["nct_id"] == "NCT00000002"
def test_nearest_sites_capped_at_five(self):
locs = [make_location(lat=42.3370 + i * 0.01, lon=-71.1061, city=f"City{i}") for i in range(8)]
studies = [make_study(locations=locs)]
result = _flatten_and_rank(studies, self.PATIENT_LAT, self.PATIENT_LON)
assert len(result[0]["nearest_sites"]) == 5
def test_site_label_contains_city_and_distance(self):
loc = make_location(lat=42.3370, lon=-71.1061, city="Boston")
studies = [make_study(locations=[loc])]
result = _flatten_and_rank(studies, self.PATIENT_LAT, self.PATIENT_LON)
label = result[0]["nearest_sites"][0]["label"]
assert "Boston" in label
assert "mi" in label
def test_central_contact_phone_fallback(self):
# Site has no phone — should fall back to central contact
loc = {
"facility": "No-Phone Clinic",
"city": "Cambridge",
"state": "MA",
"geoPoint": {"lat": 42.37, "lon": -71.10},
"contacts": [],
}
studies = [make_study(locations=[loc])]
result = _flatten_and_rank(studies, self.PATIENT_LAT, self.PATIENT_LON)
assert result[0]["nearest_sites"][0]["phone"] == "617-555-0100"
def test_location_at_zero_lat_is_silently_dropped(self):
# lat=0.0 is falsy — the `if geo.get("lat") and geo.get("lon")` guard
# skips it, so a site on the equator is treated as if it has no coordinates.
# This test documents the current (surprising) behavior.
loc = {
"facility": "Equator Clinic",
"city": "Quito",
"state": "Ecuador",
"geoPoint": {"lat": 0.0, "lon": -78.5},
"contacts": [],
}
studies = [make_study(locations=[loc])]
result = _flatten_and_rank(studies, self.PATIENT_LAT, self.PATIENT_LON)
# Because lat=0.0 is falsy the site is skipped — closest_site_miles is None
assert result[0]["closest_site_miles"] is None
assert result[0]["nearest_sites"] == []
def test_required_keys_present(self):
studies = [make_study()]
result = _flatten_and_rank(studies, self.PATIENT_LAT, self.PATIENT_LON)
row = result[0]
for key in ("nct_id", "title", "phase", "sponsor", "summary", "eligibility",
"closest_site_miles", "nearest_sites", "interventions"):
assert key in row, f"missing key: {key}"
# ---------------------------------------------------------------------------
# search_trials_api — aggFilters param construction
# ---------------------------------------------------------------------------
def _api_response(studies: list[dict], next_token: str | None = None) -> dict:
body: dict = {"studies": studies}
if next_token:
body["nextPageToken"] = next_token
return body
class TestSearchTrialsApiAggFilters:
BASE_URL = "https://clinicaltrials.gov/api/v2/studies"
def _get_params(self, httpx_mock, **kwargs) -> dict:
"""Call search_trials_api with given kwargs, return query params of the captured request."""
httpx_mock.add_response(json=_api_response([]))
search_trials_api(condition="ALS", lat=42.33, lon=-71.10, **kwargs)
return dict(httpx_mock.get_requests()[0].url.params)
def test_default_interventional_all_phases(self, httpx_mock):
params = self._get_params(httpx_mock)
assert params["aggFilters"] == "studyType:int"
def test_interventional_specific_phases(self, httpx_mock):
params = self._get_params(httpx_mock, phases=["2", "3"])
assert params["aggFilters"] == "phase:2 3"
def test_interventional_na_only_maps_to_int(self, httpx_mock):
# "na" is not a numbered phase — falls through to studyType:int
params = self._get_params(httpx_mock, phases=["na"])
assert params["aggFilters"] == "studyType:int"
def test_interventional_mixed_phases_strips_na(self, httpx_mock):
params = self._get_params(httpx_mock, phases=["2", "na"])
assert params["aggFilters"] == "phase:2"
def test_observational_study_type(self, httpx_mock):
params = self._get_params(httpx_mock, study_type="OBSERVATIONAL")
assert params["aggFilters"] == "studyType:obs"
def test_eap_no_phases(self, httpx_mock):
params = self._get_params(httpx_mock, study_type="EXPANDED_ACCESS")
assert params["aggFilters"] == "studyType:exp"
def test_eap_with_numbered_phases(self, httpx_mock):
params = self._get_params(httpx_mock, study_type="EXPANDED_ACCESS", phases=["2", "3"])
assert params["aggFilters"] == "studyType:exp,phase:2 3"
def test_eap_with_na_only_no_phase_filter(self, httpx_mock):
params = self._get_params(httpx_mock, study_type="EXPANDED_ACCESS", phases=["na"])
assert params["aggFilters"] == "studyType:exp"
def test_status_filter_recruiting_for_interventional(self, httpx_mock):
params = self._get_params(httpx_mock)
assert params["filter.overallStatus"] == "RECRUITING"
def test_status_filter_available_for_eap(self, httpx_mock):
params = self._get_params(httpx_mock, study_type="EXPANDED_ACCESS")
assert params["filter.overallStatus"] == "AVAILABLE"
def test_geo_filter_formatted(self, httpx_mock):
params = self._get_params(httpx_mock, radius_miles=50)
geo = params["filter.geo"]
# Assert structure without depending on float-to-string representation
assert geo.startswith("distance(")
assert "50mi)" in geo
assert "42.33" in geo
def test_condition_passed_through(self, httpx_mock):
params = self._get_params(httpx_mock)
assert params["query.cond"] == "ALS"
class TestSearchTrialsApiPagination:
def test_follows_next_page_token(self, httpx_mock):
page1 = _api_response([make_study("NCT00000001")], next_token="token-abc")
page2 = _api_response([make_study("NCT00000002")])
httpx_mock.add_response(json=page1)
httpx_mock.add_response(json=page2)
results = search_trials_api("ALS", lat=42.33, lon=-71.10)
assert len(results) == 2
def test_second_request_includes_page_token(self, httpx_mock):
page1 = _api_response([make_study()], next_token="token-xyz")
page2 = _api_response([])
httpx_mock.add_response(json=page1)
httpx_mock.add_response(json=page2)
search_trials_api("ALS", lat=42.33, lon=-71.10)
second_req = httpx_mock.get_requests()[1]
assert second_req.url.params["pageToken"] == "token-xyz"
def test_empty_response_returns_empty_list(self, httpx_mock):
httpx_mock.add_response(json=_api_response([]))
results = search_trials_api("ALS", lat=42.33, lon=-71.10)
assert results == []
class TestSearchTrialsApiRetry:
def test_retries_on_http_error_then_succeeds(self, httpx_mock):
httpx_mock.add_response(status_code=500)
httpx_mock.add_response(json=_api_response([make_study()]))
results = search_trials_api("ALS", lat=42.33, lon=-71.10)
assert len(results) == 1
assert len(httpx_mock.get_requests()) == 2
def test_raises_after_three_failures(self, httpx_mock):
for _ in range(3):
httpx_mock.add_response(status_code=503)
with pytest.raises(httpx.HTTPStatusError):
search_trials_api("ALS", lat=42.33, lon=-71.10)