134 lines
4.9 KiB
Python
134 lines
4.9 KiB
Python
"""Comparison tests: batched vs. unbatched runs of ``bundles_to_graph``.
|
|
|
|
Batching is driven by ``BATCH_SIZE``. The function accumulates patients into
|
|
a graph and flushes (process_references -> property_convolution ->
|
|
run_biocypher) every ``BATCH_SIZE`` patients, and once more at the end via
|
|
the ``c == n`` guard. "With batching" uses a small size; "without" uses a
|
|
size >= the patient count so there is a single flush.
|
|
|
|
We fake the collaborators: ``get_patient_everything`` returns a trivial
|
|
bundle, ``add_json_to_networkx`` adds exactly one node per patient (keyed by
|
|
the bundle name the pipeline builds, e.g. "p0_bundle"), and ``run_biocypher``
|
|
records a snapshot of the node set it is handed. That snapshot per flush is
|
|
what we compare.
|
|
"""
|
|
|
|
from unittest import mock
|
|
|
|
import pytest
|
|
|
|
import import_fhir_to_nx_diGraph as pipeline
|
|
|
|
|
|
class FakeResponse:
|
|
def __init__(self, payload):
|
|
self._payload = payload
|
|
|
|
def json(self):
|
|
return self._payload
|
|
|
|
|
|
def _run(ids, batch_size, monkeypatch):
|
|
"""Run bundles_to_graph and return the node set written at each flush."""
|
|
monkeypatch.setenv("BATCH_SIZE", str(batch_size))
|
|
monkeypatch.setattr(
|
|
pipeline, "get_patient_everything",
|
|
lambda pid: FakeResponse({"id": pid}),
|
|
)
|
|
|
|
def add_json(_bundle, name, graph):
|
|
graph.add_node(name) # one node per patient bundle
|
|
|
|
monkeypatch.setattr(
|
|
pipeline, "create_graph",
|
|
mock.Mock(add_json_to_networkx=add_json),
|
|
)
|
|
monkeypatch.setattr(pipeline, "process_references", lambda g: None)
|
|
monkeypatch.setattr(pipeline, "property_convolution", lambda g: None)
|
|
|
|
flushes = []
|
|
monkeypatch.setattr(
|
|
pipeline, "run_biocypher",
|
|
lambda graph: flushes.append(set(graph.nodes())),
|
|
)
|
|
|
|
pipeline.bundles_to_graph(ids, len(ids))
|
|
return flushes
|
|
|
|
|
|
# A big batch size relative to n => a single flush => "no batching".
|
|
NO_BATCHING = 10_000
|
|
|
|
|
|
def test_batched_and_unbatched_write_the_same_nodes(monkeypatch):
|
|
ids = [f"p{i}" for i in range(4)]
|
|
expected = {f"p{i}_bundle" for i in range(4)}
|
|
|
|
batched = _run(ids, batch_size=2, monkeypatch=monkeypatch)
|
|
unbatched = _run(ids, batch_size=NO_BATCHING, monkeypatch=monkeypatch)
|
|
|
|
# Core invariant: total written content is identical either way.
|
|
assert set().union(*batched) == expected
|
|
assert unbatched[0] == expected
|
|
assert set().union(*batched) == unbatched[0]
|
|
|
|
|
|
def test_batching_changes_flush_count_only(monkeypatch):
|
|
ids = [f"p{i}" for i in range(4)]
|
|
batched = _run(ids, batch_size=2, monkeypatch=monkeypatch)
|
|
unbatched = _run(ids, batch_size=NO_BATCHING, monkeypatch=monkeypatch)
|
|
|
|
assert len(batched) == 2 # 4 patients / batch of 2
|
|
assert len(unbatched) == 1 # single flush at the end
|
|
|
|
|
|
def test_batches_partition_patients_without_overlap(monkeypatch):
|
|
ids = [f"p{i}" for i in range(4)]
|
|
batched = _run(ids, batch_size=2, monkeypatch=monkeypatch)
|
|
|
|
# Each batch is a fresh graph, so flushes must be disjoint and together
|
|
# cover every patient exactly once.
|
|
assert batched[0].isdisjoint(batched[1])
|
|
assert sum(len(f) for f in batched) == 4
|
|
|
|
|
|
def test_uneven_batch_flushes_the_remainder(monkeypatch):
|
|
# 5 patients, batch of 2 => flushes at c=2, c=4, and c==n (the leftover).
|
|
ids = [f"p{i}" for i in range(5)]
|
|
batched = _run(ids, batch_size=2, monkeypatch=monkeypatch)
|
|
|
|
assert [len(f) for f in batched] == [2, 2, 1]
|
|
assert set().union(*batched) == {f"p{i}_bundle" for i in range(5)}
|
|
|
|
|
|
@pytest.mark.parametrize("batch_size", [1, 2, 3, 5, NO_BATCHING])
|
|
def test_no_patient_is_dropped_for_any_batch_size(batch_size, monkeypatch):
|
|
ids = [f"p{i}" for i in range(5)]
|
|
flushes = _run(ids, batch_size=batch_size, monkeypatch=monkeypatch)
|
|
|
|
# Regardless of batch size, the union of all flushes is the full set.
|
|
assert set().union(*flushes) == {f"p{i}_bundle" for i in range(5)}
|
|
|
|
|
|
def test_reference_resolution_scope_differs(monkeypatch):
|
|
# Each batched flush works on its own graph, so no single flush ever sees
|
|
# the whole graph in memory. This does NOT break correctness: cross-batch
|
|
# references are carried by dummy target nodes and resolved later by
|
|
# neo4j-admin import via id matching across all CSV part files (see
|
|
# test_dummy_node_reference_resolution.py). This test just documents the
|
|
# in-memory scope; the dummy-node tests cover the actual equivalence.
|
|
ids = [f"p{i}" for i in range(4)]
|
|
|
|
batched = _run(ids, batch_size=2, monkeypatch=monkeypatch)
|
|
unbatched = _run(ids, batch_size=NO_BATCHING, monkeypatch=monkeypatch)
|
|
|
|
assert max(len(f) for f in batched) < len(ids) # never the full graph
|
|
assert len(unbatched[0]) == len(ids) # full graph at once
|
|
|
|
|
|
def test_batch_size_zero_raises(monkeypatch):
|
|
# SHARP EDGE: BATCH_SIZE=0 makes `c % batch_size` divide by zero.
|
|
ids = [f"p{i}" for i in range(3)]
|
|
with pytest.raises(ZeroDivisionError):
|
|
_run(ids, batch_size=0, monkeypatch=monkeypatch)
|