From c1e74de512bcc03423b8fce0c70e5ae20c46a870 Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Sat, 28 Mar 2026 17:38:23 +0000 Subject: [PATCH] =?UTF-8?q?Add=20Python=20API=20tests=20to=20improve=20cov?= =?UTF-8?q?erage=20(94%=20=E2=86=92=2095%)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - matching.py: 98% → 100% (add_vestigial_root validation, AncestorMatcher optional kwargs) - config.py: 95% → 99% (field spec resolution, parsing edge cases, format optional fields) - pipeline.py: 90% → 91% (_erase_flanks no-intervals, augment_sites missing metadata) --- tests/test_config.py | 126 +++++++++++++++++++++++++++++++++++++++++ tests/test_matching.py | 44 ++++++++++++++ tests/test_pipeline.py | 63 +++++++++++++++++++++ 3 files changed, 233 insertions(+) diff --git a/tests/test_config.py b/tests/test_config.py index e976dfa3..64f6973a 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -1165,3 +1165,129 @@ def test_groups_rejected(self, tmp_path): ) with pytest.raises(ValueError, match="Unrecognised.*match.*groups"): config.Config.from_toml(_write_toml(tmp_path, toml)) + + +class TestConfigCoverageEdgeCases: + def test_resolve_field_spec_dict_with_path(self): + spec = {"path": "some/path", "field": "x"} + result = config._resolve_field_spec(spec) + assert result["path"] == "some/path" + assert result["field"] == "x" + + def test_resolve_field_spec_string(self): + result = config._resolve_field_spec("simple_field") + assert result == "simple_field" + + def test_ancestral_state_missing_path(self, tmp_path): + toml = """\ +[ancestral_state] +field = "variant_ancestral_allele" + +[[source]] +name = "cohort" +path = "samples.vcz" + +[ancestors] +path = "ancestors.vcz" +sources = ["cohort"] + +[match] +output = "out.trees" + +[match.sources.ancestors] +[match.sources.cohort] +""" + with pytest.raises(ValueError, match="missing required key"): + config.Config.from_toml(_write_toml(tmp_path, toml)) + + def test_ancestors_invalid_format(self): + raw = { + "ancestral_state": {"path": "x", "field": "y"}, + "ancestors": "not_a_table", + "match": {"sources": {"cohort": {}}, "output": "out.trees"}, + "source": [{"name": "cohort", "path": "s.vcz"}], + } + with pytest.raises(ValueError, match="must be a table"): + config._parse_ancestors(raw) + + def test_augment_sites_missing_sources(self, tmp_path): + raw = {"augment_sites": {}} + with pytest.raises(ValueError, match="missing required key.*sources"): + config._parse_augment_sites(raw) + + def test_augment_sites_sources_not_list(self, tmp_path): + raw = {"augment_sites": {"sources": "not_a_list"}} + with pytest.raises(ValueError, match="sources must be a list"): + config._parse_augment_sites(raw) + + def test_ancestor_not_in_match_sources(self): + with pytest.raises(ValueError, match="must appear in"): + config.Config( + sources={"cohort": config.Source(name="cohort", path="s.vcz")}, + ancestors=[ + config.AncestorsConfig( + name="ancestors", + path="a.vcz", + sources=["cohort"], + ) + ], + match=config.MatchConfig( + sources={"cohort": config.MatchSourceConfig()}, + output="out.trees", + ), + ancestral_state=config.AncestralState(path="ann.vcz", field="x"), + ) + + def test_match_source_simple_value(self, tmp_path): + """A match source with a bare value (not a table) gets default config.""" + toml = """\ +[ancestral_state] +path = "annotations.vcz" +field = "variant_ancestral_allele" + +[[source]] +name = "cohort" +path = "samples.vcz" + +[ancestors] +path = "ancestors.vcz" +sources = ["cohort"] + +[match] +output = "out.trees" + +[match.sources] +ancestors = true +cohort = true +""" + cfg = config.Config.from_toml(_write_toml(tmp_path, toml)) + assert isinstance(cfg.match.sources["cohort"], config.MatchSourceConfig) + + def test_config_format_optional_fields(self): + """Config.format() includes optional source fields when set.""" + cfg = config.Config( + sources={ + "cohort": config.Source( + name="cohort", + path="s.vcz", + regions="chr1:1-100", + targets="targets.bed", + ) + }, + ancestors=[ + config.AncestorsConfig( + name="ancestors", path="a.vcz", sources=["cohort"] + ) + ], + match=config.MatchConfig( + sources={ + "ancestors": config.MatchSourceConfig(), + "cohort": config.MatchSourceConfig(), + }, + output="out.trees", + ), + ancestral_state=config.AncestralState(path="ann.vcz", field="x"), + ) + text = cfg.format() + assert "regions" in text + assert "targets" in text diff --git a/tests/test_matching.py b/tests/test_matching.py index 0432da00..1028604a 100644 --- a/tests/test_matching.py +++ b/tests/test_matching.py @@ -23,6 +23,7 @@ from __future__ import annotations import numpy as np +import pytest import tskit from tsinfer import grouping, matching, vcz @@ -758,3 +759,46 @@ def test_metadata_survives_multiple_cycles(self): ) assert "sequence_intervals" in ts2.metadata assert ts2.metadata["sequence_intervals"] == [[10, 51]] + + +# --------------------------------------------------------------------------- +# TestAddVestigialRoot +# --------------------------------------------------------------------------- + + +class TestAddVestigialRoot: + def test_non_discrete_genome(self): + tables = tskit.TableCollection(sequence_length=1.5) + tables.nodes.add_row(flags=tskit.NODE_IS_SAMPLE, time=0) + ts = tables.tree_sequence() + with pytest.raises(ValueError, match="discrete genome"): + matching.add_vestigial_root(ts) + + def test_empty_tree_sequence(self): + tables = tskit.TableCollection(sequence_length=1) + ts = tables.tree_sequence() + with pytest.raises(ValueError, match="Emtpy trees"): + matching.add_vestigial_root(ts) + + +# --------------------------------------------------------------------------- +# TestAncestorMatcherWrapper +# --------------------------------------------------------------------------- + + +class TestAncestorMatcherWrapper: + def test_optional_kwargs(self): + ts = tskit.Tree.generate_balanced(4).tree_sequence + tables = ts.dump_tables() + tables.sequence_length = 2 + tables.edges.right = np.full(len(tables.edges), 2, dtype=np.float64) + tables.sites.add_row(position=1, ancestral_state="A") + tables.mutations.add_row(site=0, node=1, derived_state="T") + ts = tables.tree_sequence() + mi = matching.MatcherIndexes(ts) + r = np.full(ts.num_sites, 1e-9) + m = np.full(ts.num_sites, 0.0) + am = matching.AncestorMatcher( + mi, r, m, likelihood_threshold=1e-10, weight_by_n=False + ) + assert am is not None diff --git a/tests/test_pipeline.py b/tests/test_pipeline.py index dbf9ec66..33c29de7 100644 --- a/tests/test_pipeline.py +++ b/tests/test_pipeline.py @@ -1358,3 +1358,66 @@ def test_run_integration(self): # Should have the augmented site positions = set(out_ts.sites_position) assert 200.0 in positions + + +class TestEraseFlanks: + def _make_ts_with_json_metadata(self, metadata=None): + """Build a ts with JSON metadata schema so ts.metadata returns a dict.""" + tables = tskit.Tree.generate_balanced(4).tree_sequence.dump_tables() + tables.metadata_schema = tskit.MetadataSchema.permissive_json() + tables.metadata = metadata if metadata is not None else {} + return tables.tree_sequence() + + def test_no_intervals_returns_unchanged(self): + """_erase_flanks returns ts unchanged when metadata has no sequence_intervals.""" + ts = self._make_ts_with_json_metadata() + result = pipeline._erase_flanks(ts) + assert result.num_edges == ts.num_edges + assert result.num_nodes == ts.num_nodes + + def test_empty_metadata(self): + """_erase_flanks handles ts.metadata being empty dict.""" + ts = self._make_ts_with_json_metadata({}) + result = pipeline._erase_flanks(ts) + assert result.num_edges == ts.num_edges + + +class TestAugmentSitesEdgeCases: + def test_missing_node_metadata(self): + """augment_sites raises when sample nodes lack required metadata.""" + # Build a ts with JSON node metadata schema but nodes missing keys + tables = tskit.Tree.generate_balanced(4).tree_sequence.dump_tables() + tables.nodes.metadata_schema = tskit.MetadataSchema.permissive_json() + md = [{}] * len(tables.nodes) + tables.nodes.packset_metadata( + [tables.nodes.metadata_schema.encode_row(m) for m in md] + ) + tables.metadata_schema = tskit.MetadataSchema.permissive_json() + tables.metadata = {} + ts = tables.tree_sequence() + cfg = config.Config( + sources={"s": config.Source(name="s", path="s.vcz")}, + ancestors=[], + match=config.MatchConfig( + sources={}, output="out.trees", reference_ts="ref.trees" + ), + ancestral_state=config.AncestralState(path="ann.vcz", field="x"), + augment_sites=config.AugmentSitesConfig(sources=["s"]), + ) + with pytest.raises(ValueError, match="missing required metadata key"): + pipeline.augment_sites(ts, cfg) + + def test_no_augment_config(self): + """augment_sites returns ts unchanged when no augment config.""" + ts = tskit.Tree.generate_balanced(4).tree_sequence + cfg = config.Config( + sources={}, + ancestors=[], + match=config.MatchConfig( + sources={}, output="out.trees", reference_ts="ref.trees" + ), + ancestral_state=config.AncestralState(path="ann.vcz", field="x"), + augment_sites=None, + ) + result = pipeline.augment_sites(ts, cfg) + assert result.equals(ts)