poincare-hyper / tests /test_data_pbdb.py
DHDRL's picture
Upload 30 files
ae73c7f verified
Raw
History Blame Contribute Delete
11.7 kB
"""
Tests for src/data_pbdb.py.
Two categories:
1. Pure-function tests (occurrences_to_field, schema validation) — run
with hand-built record lists, no network needed at all.
2. Mocked-HTTP tests (fetch_occurrences, discover_schema, caching,
retries). These verify the actual pagination/retry/hard-fail
LOGIC is correct; they do not prove PBDB's real schema matches
REQUIRED_FIELDS — that must be confirmed by running
`python -m src.data_pbdb --discover` against the live API.
"""
from __future__ import annotations
import json
import os
import shutil
import sys
import tempfile
from pathlib import Path
from unittest.mock import patch, MagicMock
import pytest
import torch
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from src.data_pbdb import (
occurrences_to_field,
_validate_record_schema,
_cache_key,
fetch_occurrences,
discover_schema,
get_pbdb_dataset,
PBDBFieldDataset,
PBDBFetchError,
REQUIRED_FIELDS,
)
from src.provenance import SchemaValidationError, DataLoadError
@pytest.fixture
def tmpdir():
d = tempfile.mkdtemp()
yield d
shutil.rmtree(d, ignore_errors=True)
def make_record(max_ma, min_ma, lng, lat, genus="Testgenus"):
return {
"lng": lng, "lat": lat, "max_ma": max_ma, "min_ma": min_ma,
"genus": genus,
}
# --------------------------------------------------------------------- #
# occurrences_to_field — pure function, no network
# --------------------------------------------------------------------- #
def test_occurrences_to_field_shape():
records = [make_record(65, 60, 10, 20), make_record(50, 45, -30, -10)]
field = occurrences_to_field(records, n_time_bins=4, height=8, width=8)
assert field.shape == (4, 2, 8, 8)
def test_occurrences_to_field_places_occurrence_in_correct_cell():
records = [make_record(max_ma=100, min_ma=100, lng=0.0, lat=0.0)]
field = occurrences_to_field(records, n_time_bins=1, height=2, width=2,
age_min=0, age_max=100)
assert field[:, 0].sum().item() == 1.0 # exactly one occurrence counted
assert field[:, 1].sum().item() == 1.0 # exactly one distinct genus counted
def test_occurrences_to_field_raises_on_no_usable_records():
bad_records = [{"lng": "not-a-number", "lat": 1, "max_ma": 1, "min_ma": 1}]
with pytest.raises(SchemaValidationError) as exc_info:
occurrences_to_field(bad_records)
assert exc_info.value.outcome_code == "NO_USABLE_RECORDS"
def test_occurrences_to_field_raises_on_degenerate_age_range():
records = [make_record(50, 50, 0, 0)] * 3 # all identical age
with pytest.raises(SchemaValidationError) as exc_info:
occurrences_to_field(records, age_min=50, age_max=50)
assert exc_info.value.outcome_code == "DEGENERATE_AGE_RANGE"
def test_occurrences_to_field_never_fabricates_data_in_empty_cells():
records = [make_record(max_ma=100, min_ma=100, lng=0.0, lat=0.0)]
field = occurrences_to_field(records, n_time_bins=5, height=10, width=10,
age_min=0, age_max=100)
nonzero_cells = (field[:, 0] > 0).sum().item()
assert nonzero_cells == 1 # only the one real occurrence's cell is nonzero
# --------------------------------------------------------------------- #
# Schema validation
# --------------------------------------------------------------------- #
def test_validate_record_schema_raises_on_missing_fields():
records = [{"lng": 1, "lat": 2}] # missing max_ma, min_ma, genus
with pytest.raises(SchemaValidationError) as exc_info:
_validate_record_schema(records)
assert exc_info.value.outcome_code == "PBDB_SCHEMA_MISMATCH"
assert "max_ma" in str(exc_info.value)
def test_validate_record_schema_raises_on_empty_result():
with pytest.raises(SchemaValidationError) as exc_info:
_validate_record_schema([])
assert exc_info.value.outcome_code == "EMPTY_PBDB_RESULT"
def test_validate_record_schema_passes_with_all_required_fields():
records = [make_record(65, 60, 10, 20)]
_validate_record_schema(records) # should not raise
def test_validate_record_schema_passes_when_field_missing_only_from_first_record():
records = [
{"lng": 1, "lat": 2, "max_ma": 10, "min_ma": 5}, # genus MISSING here
{"lng": 3, "lat": 4, "max_ma": 8, "min_ma": 4, "genus": "Canis"},
]
_validate_record_schema(records) # must NOT raise
def test_validate_record_schema_still_fails_when_field_absent_from_entire_sample():
records = [{"lng": 1, "lat": 2, "max_ma": 10, "min_ma": 5} for _ in range(10)]
with pytest.raises(SchemaValidationError) as exc_info:
_validate_record_schema(records) # genus absent from all 10
assert exc_info.value.outcome_code == "PBDB_SCHEMA_MISMATCH"
def test_validate_record_schema_only_samples_first_50_records():
records = [{"lng": 1, "lat": 2, "max_ma": 10, "min_ma": 5} for _ in range(60)]
records[55]["genus"] = "Canis" # only appears past the 50-record sample window
with pytest.raises(SchemaValidationError):
_validate_record_schema(records)
# --------------------------------------------------------------------- #
# Cache key determinism
# --------------------------------------------------------------------- #
def test_cache_key_is_order_independent():
k1 = _cache_key({"base_name": "Dinosauria", "show": "full"})
k2 = _cache_key({"show": "full", "base_name": "Dinosauria"})
assert k1 == k2
def test_cache_key_differs_for_different_params():
k1 = _cache_key({"base_name": "Dinosauria"})
k2 = _cache_key({"base_name": "Mammalia"})
assert k1 != k2
# --------------------------------------------------------------------- #
# fetch_occurrences — mocked HTTP (real network blocked in this sandbox
# by paleobiodb.org's robots.txt; these verify the LOGIC, not the live
# schema — see module docstring)
# --------------------------------------------------------------------- #
def _mock_response(records, status_code=200):
resp = MagicMock()
resp.status_code = status_code
resp.json.return_value = {"records": records}
resp.text = json.dumps({"records": records})[:200]
return resp
def test_fetch_occurrences_paginates_until_short_page(tmpdir):
page1 = [make_record(65, 60, i, i, genus=f"g{i}") for i in range(3)]
page2 = [make_record(50, 45, i, i, genus=f"g{i}") for i in range(2)] # short page
call_count = {"n": 0}
def fake_get(url, params=None, headers=None, timeout=None):
call_count["n"] += 1
if params["offset"] == 0:
return _mock_response(page1)
return _mock_response(page2)
with patch("requests.get", side_effect=fake_get):
records = fetch_occurrences(
base_name="Testtaxon", max_records=100, page_size=3,
cache_dir=os.path.join(tmpdir, "cache"),
)
assert len(records) == 5 # 3 + 2, stopped because page2 was short
assert call_count["n"] == 2
def test_fetch_occurrences_stops_at_max_records(tmpdir):
full_page = [make_record(65, 60, i, i, genus=f"g{i}") for i in range(5)]
def fake_get(url, params=None, headers=None, timeout=None):
return _mock_response(full_page[: params["limit"]])
with patch("requests.get", side_effect=fake_get):
records = fetch_occurrences(
base_name="Testtaxon", max_records=7, page_size=5,
cache_dir=os.path.join(tmpdir, "cache"),
)
assert len(records) <= 7
def test_fetch_occurrences_hard_fails_after_retries_exhausted(tmpdir):
def fake_get(url, params=None, headers=None, timeout=None):
return _mock_response([], status_code=500)
with patch("requests.get", side_effect=fake_get), \
patch("time.sleep", return_value=None): # skip real backoff delays in test
with pytest.raises(PBDBFetchError) as exc_info:
fetch_occurrences(base_name="Testtaxon", cache_dir=os.path.join(tmpdir, "cache"))
assert exc_info.value.outcome_code == "PBDB_HTTP_ERROR"
def test_fetch_occurrences_hard_fails_on_schema_mismatch(tmpdir):
bad_records = [{"lng": 1, "lat": 2}] # missing max_ma/min_ma/genus
def fake_get(url, params=None, headers=None, timeout=None):
return _mock_response(bad_records)
with patch("requests.get", side_effect=fake_get):
with pytest.raises(SchemaValidationError) as exc_info:
fetch_occurrences(base_name="Testtaxon", cache_dir=os.path.join(tmpdir, "cache"))
assert exc_info.value.outcome_code == "PBDB_SCHEMA_MISMATCH"
def test_fetch_occurrences_uses_cache_on_second_call(tmpdir):
records = [make_record(65, 60, 1, 1, genus="g1")]
call_count = {"n": 0}
def fake_get(url, params=None, headers=None, timeout=None):
call_count["n"] += 1
return _mock_response(records)
cache_dir = os.path.join(tmpdir, "cache")
with patch("requests.get", side_effect=fake_get):
r1 = fetch_occurrences(base_name="Testtaxon", cache_dir=cache_dir)
r2 = fetch_occurrences(base_name="Testtaxon", cache_dir=cache_dir)
assert r1 == r2
assert call_count["n"] == 1 # second call hit the cache, not the network
def test_discover_schema_reports_missing_fields(tmpdir):
incomplete_records = [{"lng": 1, "lat": 2, "genus": "g1"}] # no max_ma/min_ma
def fake_get(url, params=None, headers=None, timeout=None):
return _mock_response(incomplete_records)
with patch("requests.get", side_effect=fake_get):
result = discover_schema(base_name="Testtaxon", limit=1)
assert "max_ma" in result["required_fields_missing"]
assert "genus" in result["required_fields_present"]
def test_discover_schema_hard_fails_on_empty_result(tmpdir):
def fake_get(url, params=None, headers=None, timeout=None):
return _mock_response([])
with patch("requests.get", side_effect=fake_get):
with pytest.raises(PBDBFetchError) as exc_info:
discover_schema(base_name="Nonexistenttaxon")
assert exc_info.value.outcome_code == "PBDB_EMPTY_RESPONSE"
# --------------------------------------------------------------------- #
# End-to-end (mocked): PBDBFieldDataset / get_pbdb_dataset
# --------------------------------------------------------------------- #
def test_pbdb_field_dataset_one_trajectory_per_taxon_group(tmpdir):
records = [make_record(max_ma=100 - i * 10, min_ma=95 - i * 10,
lng=i, lat=i, genus=f"g{i}") for i in range(4)]
def fake_get(url, params=None, headers=None, timeout=None):
return _mock_response(records)
with patch("requests.get", side_effect=fake_get):
ds = PBDBFieldDataset(
taxon_groups=["TaxonA", "TaxonB", "TaxonC"],
n_time_bins=4, height=8, width=8,
cache_dir=os.path.join(tmpdir, "cache"),
)
assert len(ds) == 3
item = ds[0]
assert item["fields"].shape == (4, 2, 8, 8)
assert item["taxon_group"] == "TaxonA"
def test_get_pbdb_dataset_provenance_is_real_pbdb(tmpdir):
records = [make_record(max_ma=100 - i * 10, min_ma=95 - i * 10,
lng=i, lat=i, genus=f"g{i}") for i in range(4)]
def fake_get(url, params=None, headers=None, timeout=None):
return _mock_response(records)
with patch("requests.get", side_effect=fake_get):
ds, provenance = get_pbdb_dataset(
taxon_groups=["TaxonA"], cache_dir=os.path.join(tmpdir, "cache"),
)
assert provenance == "REAL_PBDB"
assert len(ds) == 1
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v"]))